diff --git a/app/dispatcher/fakednssniffer.go b/app/dispatcher/fakednssniffer.go index bed90877b..9476288c7 100644 --- a/app/dispatcher/fakednssniffer.go +++ b/app/dispatcher/fakednssniffer.go @@ -23,7 +23,7 @@ func newFakeDNSSniffer(ctx context.Context) (protocolSnifferWithMetadata, error) } 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{protocolSniffer: func(ctx context.Context, bytes []byte) (SniffResult, error) { diff --git a/app/dns/config.go b/app/dns/config.go index 1f9ac0153..8b7160ba7 100644 --- a/app/dns/config.go +++ b/app/dns/config.go @@ -28,7 +28,7 @@ func toNetIP(addrs []net.Address) ([]net.IP, error) { if addr.Family().IsIP() { ips = append(ips, addr.IP()) } 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 diff --git a/app/dns/dns.go b/app/dns/dns.go index 6ea28ecf2..9943da09d 100644 --- a/app/dns/dns.go +++ b/app/dns/dns.go @@ -225,6 +225,28 @@ func (s *DNS) IsOwnLink(ctx context.Context) bool { 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. func (s *DNS) LookupIP(domain string, option dns.IPOption) ([]net.IP, uint32, error) { // Normalize the FQDN form query diff --git a/app/dns/dns_internal_test.go b/app/dns/dns_internal_test.go new file mode 100644 index 000000000..7614efd8b --- /dev/null +++ b/app/dns/dns_internal_test.go @@ -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) + } + }) + } +} diff --git a/app/dns/dnscommon.go b/app/dns/dnscommon.go index 5ffae51be..10597d5b3 100644 --- a/app/dns/dnscommon.go +++ b/app/dns/dnscommon.go @@ -188,10 +188,10 @@ func parseResponse(payload []byte) (*IPRecord, error) { var parser dnsmessage.Parser h, err := parser.Start(payload) 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 { - 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() diff --git a/app/dns/fakedns/fake.go b/app/dns/fakedns/fake.go index 1539e5130..e00f0ec6a 100644 --- a/app/dns/fakedns/fake.go +++ b/app/dns/fakedns/fake.go @@ -58,7 +58,7 @@ func NewFakeDNSHolder() (*Holder, error) { var err error 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) if err != nil { @@ -80,13 +80,13 @@ func (fkdns *Holder) initialize(ipPoolCidr string, lruSize int) error { var err error 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() rooms := bits - ones 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.ipRange = ipRange diff --git a/app/dns/nameserver.go b/app/dns/nameserver.go index 27ce9a540..5112f3581 100644 --- a/app/dns/nameserver.go +++ b/app/dns/nameserver.go @@ -85,7 +85,7 @@ func NewServer(ctx context.Context, dest net.Destination, dispatcher routing.Dis if dest.Network == net.Network_UDP { // UDP classic DNS mode 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. @@ -103,7 +103,7 @@ func NewClient( // Create a new server for each client for now server, err := NewServer(ctx, ns.Address.AsDestination(), dispatcher, disableCache, serveStale, serveExpiredTTL, clientIP) 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) @@ -114,7 +114,7 @@ func NewClient( if len(ns.ExpectedIp) > 0 { expectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.ExpectedIp) 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) } } @@ -123,7 +123,7 @@ func NewClient( if len(ns.UnexpectedIp) > 0 { unexpectedMatcher, err = geodata.IPReg.BuildIPMatcher(ns.UnexpectedIp) 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) } } diff --git a/app/dns/nameserver_fakedns.go b/app/dns/nameserver_fakedns.go index bed11bd60..f2127c9e8 100644 --- a/app/dns/nameserver_fakedns.go +++ b/app/dns/nameserver_fakedns.go @@ -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) { 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 @@ -39,7 +39,7 @@ func (f *FakeDNSServer) QueryIP(ctx context.Context, domain string, opt dns.IPOp netIP, err := toNetIP(ips) 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) diff --git a/app/log/log.go b/app/log/log.go index 553862b17..4d923365b 100644 --- a/app/log/log.go +++ b/app/log/log.go @@ -89,10 +89,10 @@ func (g *Instance) startInternal() error { g.active = true 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 { - return errors.New("failed to initialize error logger").Base(err).AtWarning() + return errors.New("failed to initialize error logger").Base(err) } 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(). func (g *Instance) Close() error { errors.LogDebug(context.Background(), "Logger closing") diff --git a/app/proxyman/inbound/always.go b/app/proxyman/inbound/always.go index e0d6c0925..70edbef6c 100644 --- a/app/proxyman/inbound/always.go +++ b/app/proxyman/inbound/always.go @@ -66,7 +66,7 @@ func NewAlwaysOnInboundHandler(ctx context.Context, tag string, receiverConfig * } mss, err := internet.ToMemoryStreamConfig(receiverConfig.StreamSettings) 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}) diff --git a/app/proxyman/inbound/inbound.go b/app/proxyman/inbound/inbound.go index cec93dd48..411de8598 100644 --- a/app/proxyman/inbound/inbound.go +++ b/app/proxyman/inbound/inbound.go @@ -165,7 +165,7 @@ func NewHandler(ctx context.Context, config *core.InboundHandlerConfig) (inbound receiverSettings, ok := rawReceiverSettings.(*proxyman.ReceiverConfig) if !ok { - return nil, errors.New("not a ReceiverConfig").AtError() + return nil, errors.New("not a ReceiverConfig") } streamSettings := receiverSettings.StreamSettings diff --git a/app/proxyman/inbound/worker.go b/app/proxyman/inbound/worker.go index b91c9ec66..4a3bb6d55 100644 --- a/app/proxyman/inbound/worker.go +++ b/app/proxyman/inbound/worker.go @@ -142,7 +142,7 @@ func (w *tcpWorker) Start() error { go w.callback(conn) }) 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 return nil @@ -528,7 +528,7 @@ func (w *dsWorker) Start() error { go w.callback(conn) }) 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 return nil diff --git a/app/proxyman/outbound/handler.go b/app/proxyman/outbound/handler.go index 136796f2c..fe43c429f 100644 --- a/app/proxyman/outbound/handler.go +++ b/app/proxyman/outbound/handler.go @@ -87,7 +87,7 @@ func NewHandler(ctx context.Context, config *core.OutboundHandlerConfig) (outbou h.senderSettings = s mss, err := internet.ToMemoryStreamConfig(s.StreamSettings) 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 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 { switch h.udp443 { case "reject": - test(errors.New("XUDP rejected UDP/443 traffic").AtInfo()) + test(errors.New("XUDP rejected UDP/443 traffic")) return case "skip": goto out diff --git a/app/reverse/portal.go b/app/reverse/portal.go index 7e3f2cafd..0e595983f 100644 --- a/app/reverse/portal.go +++ b/app/reverse/portal.go @@ -68,13 +68,13 @@ func (p *Portal) HandleConnection(ctx context.Context, link *transport.Link) err outbounds := session.OutboundsFromContext(ctx) ob := outbounds[len(outbounds)-1] if ob == nil { - return errors.New("outbound metadata not found").AtError() + return errors.New("outbound metadata not found") } if isDomain(ob.Target, p.domain) { muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{}) 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) diff --git a/app/router/config.go b/app/router/config.go index 036b24ef5..c20e120d6 100644 --- a/app/router/config.go +++ b/app/router/config.go @@ -115,7 +115,7 @@ func (rr *RoutingRule) BuildCondition() (Condition, error) { } 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 @@ -145,7 +145,7 @@ func (br *BalancingRule) Build(ohm outbound.Manager, dispatcher routing.Dispatch } s, ok := i.(*StrategyLeastLoadConfig) if !ok { - return nil, errors.New("not a StrategyLeastLoadConfig").AtError() + return nil, errors.New("not a StrategyLeastLoadConfig") } leastLoadStrategy := NewLeastLoadStrategy(s) return &Balancer{ diff --git a/common/errors/errors.go b/common/errors/errors.go index 7a35f2543..674714e66 100644 --- a/common/errors/errors.go +++ b/common/errors/errors.go @@ -18,17 +18,12 @@ type hasInnerError interface { Unwrap() error } -type hasSeverity interface { - Severity() log.Severity -} - // Error is an error object with underlying error. type Error struct { - prefix []interface{} - message []interface{} - caller string - inner error - severity log.Severity + prefix []interface{} + message []interface{} + caller string + inner error } // Error implements error.Error(). @@ -69,46 +64,6 @@ func (err *Error) Base(e error) *Error { 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. func (err *Error) String() string { return err.Error() @@ -132,9 +87,8 @@ func New(msg ...interface{}) *Error { details = details[:i] } return &Error{ - message: msg, - severity: log.Severity_Info, - caller: details, + message: msg, + 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{}) { + if log.GetSeverity() < severity { + return + } pc, _, _, _ := runtime.Caller(2) details := runtime.FuncForPC(pc).Name() if len(details) >= trim { @@ -181,10 +138,9 @@ func doLog(ctx context.Context, inner error, severity log.Severity, msg ...inter details = details[:i] } err := &Error{ - message: msg, - severity: severity, - caller: details, - inner: inner, + message: msg, + caller: details, + inner: inner, } if ctx != nil && ctx != context.Background() { 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{ - Severity: GetSeverity(err), + Severity: severity, Content: err, }) } @@ -217,11 +173,3 @@ L: } 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 -} diff --git a/common/errors/errors_test.go b/common/errors/errors_test.go index 3a1cb134f..a11eaac27 100644 --- a/common/errors/errors_test.go +++ b/common/errors/errors_test.go @@ -7,30 +7,21 @@ import ( "github.com/google/go-cmp/cmp" . "github.com/xtls/xray-core/common/errors" - "github.com/xtls/xray-core/common/log" ) func TestError(t *testing.T) { err := New("TestError") - if v := GetSeverity(err); v != log.Severity_Info { - t.Error("severity: ", v) + if v := err.Error(); !strings.Contains(v, "TestError") { + t.Error("error: ", v) } err = New("TestError2").Base(io.EOF) - if v := GetSeverity(err); v != log.Severity_Info { - t.Error("severity: ", v) + if v := err.Error(); !strings.Contains(v, "EOF") { + t.Error("error: ", v) } - err = New("TestError3").Base(io.EOF).AtWarning() - if v := GetSeverity(err); v != log.Severity_Warning { - 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) - } + err = New("TestError3").Base(io.EOF) + err = New("TestError4").Base(err) if v := err.Error(); !strings.Contains(v, "EOF") { t.Error("error: ", v) } diff --git a/common/geodata/domain_matcher.go b/common/geodata/domain_matcher.go index a4d3564e6..e7227f58d 100644 --- a/common/geodata/domain_matcher.go +++ b/common/geodata/domain_matcher.go @@ -82,19 +82,10 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch } g.Add(m, uint32(i)) 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 { return nil, err } - for j, d := range domains { - domains[j] = nil // peak mem - m, err := parseDomain(d) - if err != nil { - errors.LogError(context.Background(), "ignore invalid geosite entry in ", v.Geosite.File, ":", v.Geosite.Code, " at index ", j, ", ", err) - continue - } - g.Add(m, uint32(i)) - } default: panic("unknown domain rule type") } @@ -108,12 +99,12 @@ func (f *MphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatch return g, nil } -type CompactDomainMatcherFactory struct { +type CompactMphDomainMatcherFactory struct { 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 f.Lock() @@ -125,33 +116,23 @@ func (f *CompactDomainMatcherFactory) getOrCreateFrom(rule *GeoSiteRule) (strmat } errors.LogDebug(context.Background(), "geodata geosite matcher cache MISS ", key) - s := strmatcher.NewLinearAnyMatcher() - domains, err := loadSiteWithAttrs(rule.File, rule.Code, rule.Attrs) - if err != nil { + s := strmatcher.NewMphValueMatcher() + if err := loadSiteMatchers(rule, func(m strmatcher.Matcher) { s.Add(m, 0) }); err != nil { return nil, err } - for i, d := range domains { - domains[i] = nil // peak mem - m, err := parseDomain(d) - if err != nil { - errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err) - continue - } - s.Add(m) + if err := s.Build(); err != nil { + return nil, err } f.shared.Store(key, s) - return s, err + return s, nil } // BuildMatcher implements DomainMatcherFactory. -func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) { +func (f *CompactMphDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainMatcher, error) { if len(rules) == 0 { return nil, errors.New("empty domain rule list") } - compact := &CompactDomainMatcher{ - matchers: make([]strmatcher.MatcherSet, 0, len(rules)), - values: make([]uint32, 0, len(rules)), - } + compact := new(CompactMphDomainMatcher) for i, r := range rules { switch v := r.Value.(type) { case *DomainRule_Custom: @@ -168,8 +149,7 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM if err != nil { return nil, err } - compact.matchers = append(compact.matchers, m) - compact.values = append(compact.values, uint32(i)) + compact.combiner.Add(m, uint32(i)) default: panic("unknown domain rule type") } @@ -177,37 +157,40 @@ func (f *CompactDomainMatcherFactory) BuildMatcher(rules []*DomainRule) (DomainM return compact, nil } -type CompactDomainMatcher struct { +type CompactMphDomainMatcher struct { custom strmatcher.ValueMatcher - matchers []strmatcher.MatcherSet - values []uint32 + combiner strmatcher.MphValueMatcherCombiner } // Match implements DomainMatcher. -func (c *CompactDomainMatcher) Match(input string) []uint32 { - var result []uint32 +func (c *CompactMphDomainMatcher) Match(input string) []uint32 { + result := c.combiner.Match(input) if c.custom != nil { - result = append(result, c.custom.Match(input)...) - } - for i, m := range c.matchers { - if m.MatchAny(input) { - result = append(result, c.values[i]) - } + result = append(c.custom.Match(input), result...) } return result } // 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) { return true } - for _, m := range c.matchers { - if m.MatchAny(input) { - return true + return c.combiner.MatchAny(input) +} + +// loadSiteMatchers calls add with a matcher for every domain of the geosite rule and logs the invalid ones. +func loadSiteMatchers(rule *GeoSiteRule, add func(strmatcher.Matcher)) error { + i := 0 + return loadSite(rule.File, rule.Code, rule.Attrs, func(t Domain_Type, value []byte) { + m, err := parseDomain(&Domain{Type: t, Value: string(value)}) + if err != nil { + errors.LogError(context.Background(), "ignore invalid geosite entry in ", rule.File, ":", rule.Code, " at index ", i, ", ", err) + } else { + add(m) } - } - return false + i++ + }) } func parseDomain(d *Domain) (strmatcher.Matcher, error) { @@ -231,7 +214,7 @@ func parseDomain(d *Domain) (strmatcher.Matcher, error) { func newDomainMatcherFactory() DomainMatcherFactory { switch runtime.GOOS { case "ios", "android": - return &CompactDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.LinearAnyMatcher]()} + return &CompactMphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()} default: return &MphDomainMatcherFactory{shared: utils.NewWeakCacheMap[string, strmatcher.MphValueMatcher]()} } diff --git a/common/geodata/domain_matcher_test.go b/common/geodata/domain_matcher_test.go index 0c0c5080c..dc3558020 100644 --- a/common/geodata/domain_matcher_test.go +++ b/common/geodata/domain_matcher_test.go @@ -4,6 +4,7 @@ import ( "path/filepath" "reflect" "slices" + "sync" "testing" "github.com/xtls/xray-core/common/geodata/strmatcher" @@ -11,7 +12,7 @@ import ( ) 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{ {Value: &DomainRule_Custom{Custom: &Domain{Type: Domain_Full, 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) { 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{ {Value: &DomainRule_Geosite{Geosite: &GeoSiteRule{File: DefaultGeoSiteDat, Code: "CN"}}}, {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}) } } + +// 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() + }) + } +} diff --git a/common/geodata/geodat_loader.go b/common/geodata/geodat_loader.go index 0e12aa28b..81d9e7e90 100644 --- a/common/geodata/geodat_loader.go +++ b/common/geodata/geodat_loader.go @@ -5,11 +5,14 @@ import ( "bytes" "io" "runtime" + "slices" "strings" + "unicode/utf8" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/platform/filesystem" + "google.golang.org/protobuf/encoding/protowire" "google.golang.org/protobuf/proto" ) @@ -52,17 +55,56 @@ func loadIP(file, code string) ([]*CIDR, error) { return geoip.Cidr, nil } -func loadSite(file, code string) ([]*Domain, error) { - bs, err := loadFile(file, code) +// loadSite calls fn, in file order, with the type and value of every domain of the geosite code +// that has all the "@"-separated attrs. It decodes the entry while reading the file instead of +// unmarshalling it into a []*Domain, so value is only valid during fn. +func loadSite(file, code, attrs string, fn func(Domain_Type, []byte)) error { + runtime.GC() // peak mem + r, err := filesystem.OpenAsset(file) if err != nil { - return nil, err + return errors.New("failed to open ", file).Base(err) } - defer runtime.GC() // peak mem - var geosite GeoSite - if err := proto.Unmarshal(bs, &geosite); err != nil { - return nil, errors.New("error unmarshal Site in ", file, ":", code).Base(err) + defer r.Close() + br := bufio.NewReaderSize(r, 64*1024) + n, err := seek(br, []byte(code)) + if err != nil { + return errors.New("failed to load code ", code, " from ", file).Base(err) } - 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) { @@ -82,68 +124,63 @@ func decodeVarint(br *bufio.Reader) (uint64, error) { } func find(r io.Reader, code []byte, readBody bool) ([]byte, error) { + br := bufio.NewReaderSize(r, 64*1024) + bodyL, err := seek(br, code) + if err != nil || !readBody { + return nil, err + } + out := make([]byte, bodyL) + if _, err := io.ReadFull(br, out); err != nil { + return nil, err + } + return out, nil +} + +// seek advances br to the body of the entry for code and returns the body length. +func seek(br *bufio.Reader, code []byte) (int, error) { codeL := len(code) if codeL == 0 { - return 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 - prefixBuf := make([]byte, need) for { if _, err := br.ReadByte(); err != nil { - return nil, err + return 0, err } x, err := decodeVarint(br) if err != nil { - return nil, err + return 0, err } bodyL := int(x) if bodyL <= 0 { - return nil, errors.New("invalid body length: ", bodyL) + return 0, errors.New("invalid body length: ", bodyL) } - prefixL := bodyL - if prefixL > need { - prefixL = need - } - prefix := prefixBuf[:prefixL] - if _, err := io.ReadFull(br, prefix); err != nil { - return nil, err - } - - match := false - if bodyL >= need { - if int(prefix[1]) == codeL && bytes.Equal(prefix[2:need], code) { - if !readBody { - return nil, nil - } - match = true + // Peek no more than the buffer holds: a code longer than the buffer cannot match a single + // length byte anyway, so a short peek only skips it, as base find (io.ReadFull) does. + prefix, err := br.Peek(min(bodyL, need, br.Size())) + if err != nil { + if err == io.EOF && len(prefix) > 0 { + err = io.ErrUnexpectedEOF // as io.ReadFull } + return 0, err } - - remain := bodyL - prefixL - if match { - out := make([]byte, bodyL) - copy(out, prefix) - if remain > 0 { - if _, err := io.ReadFull(br, out[prefixL:]); err != nil { - return nil, err - } - } - return out, nil + if bodyL >= need && len(prefix) >= need && int(prefix[1]) == codeL && bytes.Equal(prefix[2:], code) { + return bodyL, nil } - - if remain > 0 { - if _, err := br.Discard(remain); err != nil { - return nil, err - } + if _, err := br.Discard(bodyL); err != nil { + return 0, err } } } +// AttributeMatcher, HasAttrMatcher, AllAttrsMatcher and NewAllAttrsMatcher are the exported +// attribute helpers that have been part of this package's API since #5814. The streaming loader +// above filters attributes itself without building a *Domain, so it does not use them, but they +// are kept for external callers. Their behaviour is unchanged. + type AttributeMatcher interface { Match(*Domain) bool } @@ -185,23 +222,137 @@ func NewAllAttrsMatcher(attrs string) AttributeMatcher { return m } -func loadSiteWithAttrs(file, code, attrs string) ([]*Domain, error) { - domains, err := loadSite(file, code) - if err != nil { - return nil, err - } +var errInvalidUTF8 = errors.New("string field contains invalid UTF-8") - matcher := NewAllAttrsMatcher(attrs) - if matcher == nil { - return domains, nil - } +type siteDecoder struct { + want []string + has []bool + fn func(Domain_Type, []byte) +} - filtered := make([]*Domain, 0, len(domains)) - for _, d := range domains { - if matcher.Match(d) { - filtered = append(filtered, d) +func newSiteDecoder(attrs string, fn func(Domain_Type, []byte)) *siteDecoder { + d := &siteDecoder{fn: fn} + if attrs != "" { + d.want = strings.Split(attrs, "@") + d.has = make([]bool, len(d.want)) + } + 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 } diff --git a/common/geodata/geodat_loader_test.go b/common/geodata/geodat_loader_test.go new file mode 100644 index 000000000..b6d400294 --- /dev/null +++ b/common/geodata/geodat_loader_test.go @@ -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") + } +} diff --git a/common/geodata/strmatcher/benchmark_matchers_test.go b/common/geodata/strmatcher/benchmark_matchers_test.go index 9e00c816c..2cc82677d 100644 --- a/common/geodata/strmatcher/benchmark_matchers_test.go +++ b/common/geodata/strmatcher/benchmark_matchers_test.go @@ -1,6 +1,7 @@ package strmatcher_test import ( + "regexp" "strconv" "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 func benchmarkMatcherType(b *testing.B, t Type, ctor func() MatcherGroup) { diff --git a/common/geodata/strmatcher/indexmatcher_mph.go b/common/geodata/strmatcher/indexmatcher_mph.go index b23f83763..23e4d71ae 100644 --- a/common/geodata/strmatcher/indexmatcher_mph.go +++ b/common/geodata/strmatcher/indexmatcher_mph.go @@ -52,7 +52,9 @@ func (g *MphIndexMatcher) Add(matcher Matcher) uint32 { func (g *MphIndexMatcher) Build() error { if g.mph != nil { runtime.GC() // peak mem - g.mph.Build() + if err := g.mph.Build(); err != nil { + return err + } } runtime.GC() // peak mem if g.ac != nil { @@ -64,23 +66,17 @@ func (g *MphIndexMatcher) Build() error { // Match implements IndexMatcher.Match. func (g *MphIndexMatcher) Match(input string) []uint32 { - result := make([][]uint32, 0, 5) + var result []uint32 if g.mph != nil { - if matches := g.mph.Match(input); len(matches) > 0 { - result = append(result, matches) - } + result = g.mph.Match(input) // a new slice, returned without another copy } if g.ac != nil { - if matches := g.ac.Match(input); len(matches) > 0 { - result = append(result, matches) - } + result = append(result, g.ac.Match(input)...) } if g.regex != nil { - if matches := g.regex.Match(input); len(matches) > 0 { - result = append(result, matches) - } + result = append(result, g.regex.Match(input)...) } - return CompositeMatches(result) + return result } // MatchAny implements IndexMatcher.MatchAny. diff --git a/common/geodata/strmatcher/indexmatcher_mph_test.go b/common/geodata/strmatcher/indexmatcher_mph_test.go index 2e1c70dd2..efc747a71 100644 --- a/common/geodata/strmatcher/indexmatcher_mph_test.go +++ b/common/geodata/strmatcher/indexmatcher_mph_test.go @@ -78,6 +78,10 @@ func TestMphIndexMatcher(t *testing.T) { Input: "example.com", Output: []uint32{10, 4}, }, + { + Input: "apis.org", + Output: []uint32{2, 6}, + }, } matcherGroup := NewMphIndexMatcher() for _, rule := range rules { @@ -87,8 +91,13 @@ func TestMphIndexMatcher(t *testing.T) { } matcherGroup.Build() 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) } + 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) + } } } diff --git a/common/geodata/strmatcher/matchergroup_mph.go b/common/geodata/strmatcher/matchergroup_mph.go index ebf5ff7f3..2db3cc9d6 100644 --- a/common/geodata/strmatcher/matchergroup_mph.go +++ b/common/geodata/strmatcher/matchergroup_mph.go @@ -1,198 +1,440 @@ package strmatcher import ( - "math/bits" - "runtime" - "sort" + "bytes" + "cmp" + "encoding/binary" + "errors" + "math" + "slices" "strings" "unsafe" ) -// PrimeRK is the prime base used in Rabin-Karp algorithm. -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 performs as distinct hash functions. -func MemHash(seed uint32, input string) uint32 { - return uint32(strhash(unsafe.Pointer(&input), uintptr(seed))) // nosemgrep -} - +// Flags of a level1 slot, stored above the record offset. 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 { - rollingHash uint32 - matchers [mphMatchTypeCount][]uint32 +// Kinds of an added pattern, indexes of mphKinds. +const ( + mphKindFull = iota + mphKindParent + mphKindDomain +) + +// mphKinds are the slot flags in the order Match reports their values. +var mphKinds = [...]uint32{mphFull, mphParent, mphDomain} + +// mphMultipliers are odd multipliers for the suffix hash. Build moves to the next one if two patterns collide. +var mphMultipliers = [...]uint64{0x9e3779b97f4a7c15, 0xc2b2ae3d27d4eb4f, 0x165667b19e3779f9, 0x27d4eb2f165667c5} + +var ( + errMphCollision = errors.New("strmatcher: suffix hash collision in MphMatcherGroup") + errMphBuilt = errors.New("strmatcher: MphMatcherGroup is already built") +) + +type mphEntry struct { + off uint32 // pattern start in buf + value uint32 + n uint32 // pattern length + kind uint8 } -// MphMatcherGroup is an implementation of MatcherGroup. -// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher. +// MphMatcherGroup is an implementation of MatcherGroup for Full and Domain matchers. +// Each distinct pattern is stored once as a record in arena: its length (255 means a uvarint length follows), +// its bytes and, if the group holds more than one distinct value, its values. A minimal perfect hash table +// built with hash, displace and compress (http://cmph.sourceforge.net/papers/esa09.pdf) maps a pattern to its +// record. Patterns are hashed from the right, so one pass over the input hashes all its parent domains. type MphMatcherGroup struct { - rules []string // RuleIdx -> pattern string, index 0 reserved for failed lookup - values [][]uint32 // RuleIdx -> registered matcher values for the pattern (Full Matcher takes precedence) - level0 []uint32 // RollingHash & Mask -> seed for Memhash - level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0) - level1 []uint32 // Memhash & Mask -> stored index for rules - level1Mask uint32 // Mask for restricting Memhash to 0 ~ len(level1) - ruleInfos *map[string]mphRuleInfo + arena string + level0 []uint16 // bucket -> seed + level1 []uint32 // slot -> flags | record offset + fp []uint8 // slot -> low byte of its pattern's hash, rejects most misses without reading arena + n0, n1 uint32 + mul uint64 // multiplier of the suffix hash + single uint32 // the only value if !multi + multi bool + + buf []byte // build only, patterns in Add order + entries []mphEntry } func NewMphMatcherGroup() *MphMatcherGroup { - return &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 - } + return new(MphMatcherGroup) } // AddFullMatcher implements MatcherGroupForFull. func (g *MphMatcherGroup) AddFullMatcher(matcher FullMatcher, value uint32) { - pattern := strings.ToLower(matcher.Pattern()) - g.addPattern(0, "", pattern, matcher.Type(), value) + g.add(matcher.Pattern(), mphKindFull, value) } // AddDomainMatcher implements MatcherGroupForDomain. func (g *MphMatcherGroup) AddDomainMatcher(matcher DomainMatcher, value uint32) { - pattern := strings.ToLower(matcher.Pattern()) - hash := g.addPattern(0, "", pattern, matcher.Type(), value) // For full domain match - g.addPattern(hash, pattern, ".", matcher.Type(), value) // For partial domain match + g.add(matcher.Pattern(), mphKindDomain, value) } -func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pattern string, matcherType Type, value uint32) uint32 { - fullPattern := pattern + suffixPattern - info, found := (*g.ruleInfos)[fullPattern] - if !found { - info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)} - g.rules = append(g.rules, fullPattern) - g.values = append(g.values, nil) +func (g *MphMatcherGroup) add(pattern string, kind uint8, value uint32) { + if g.arena != "" { + panic(errMphBuilt) + } + pattern = strings.ToLower(pattern) + off := uint32(len(g.buf)) + g.buf = append(g.buf, pattern...) + g.entries = append(g.entries, mphEntry{off: off, value: value, n: uint32(len(pattern)), kind: kind}) + if len(pattern) > 0 && pattern[0] == '.' { + // ".x" has always matched "*.x" as well, so it also gets a parent-only record for "x" + g.entries = append(g.entries, mphEntry{off: off + 1, value: value, n: uint32(len(pattern) - 1), kind: mphKindParent}) } - 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. -// Algorithm used: Hash, displace, and compress. See http://cmph.sourceforge.net/papers/esa09.pdf +func (g *MphMatcherGroup) key(i uint32) []byte { + e := &g.entries[i] + return g.buf[e.off : e.off+e.n] +} + +// Build builds the hash table. It must be called once, after the last Add. func (g *MphMatcherGroup) Build() error { - ruleCount := len(*g.ruleInfos) - g.level0 = make([]uint32, nextPow2(ruleCount/4)) - g.level0Mask = uint32(len(g.level0) - 1) - g.level1 = make([]uint32, nextPow2(ruleCount)) - g.level1Mask = uint32(len(g.level1) - 1) - - // 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 + if g.arena != "" { + return errMphBuilt } - g.ruleInfos = nil // Set ruleInfos nil to release memory - runtime.GC() // peak mem - - // Sort buckets in descending order with respect to each bucket's size - bucketIdxs := make([]int, len(buckets)) - for bucketIdx := range buckets { - bucketIdxs[bucketIdx] = bucketIdx + if uint64(len(g.buf)) > math.MaxUint32 { + return errors.New("too many rules for MphMatcherGroup") } - 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 - occupied := make([]bool, len(g.level1)) // Whether a second-level hash has been already used - hashedBucket := make([]uint32, 0, 4) // Second-level hashes for each rule in a specific bucket - for _, bucketIdx := range bucketIdxs { - bucket := buckets[bucketIdx] - hashedBucket = hashedBucket[:0] - seed := uint32(0) - for len(hashedBucket) != len(bucket) { - for _, ruleIdx := range bucket { - memHash := MemHash(seed, g.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) +// writeRecords writes one record per distinct pattern to arena and returns flags | offset of each. +func (g *MphMatcherGroup) writeRecords() []uint32 { + g.multi = false + if len(g.entries) > 0 { + g.single = g.entries[0].value + for _, e := range g.entries { + if e.value != g.single { + g.multi = true + break } } - 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 } -// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found. -func (g *MphMatcherGroup) Lookup(rollingHash uint32, input string) uint32 { - i0 := rollingHash & g.level0Mask - seed := g.level0[i0] - i1 := MemHash(seed, input) & g.level1Mask - if n := g.level1[i1]; g.rules[n] == input { - return n +// mphHash is the suffix hash of s, taken from the right: the hash of s[i:] is the state after reading s[i]. +func mphHash(mul uint64, s string) uint64 { + h := uint64(0) + for i := len(s) - 1; i >= 0; i-- { + h = h*mul + uint64(s[i]) + } + return h +} + +// 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 } -// Match implements MatcherGroup.Match. -func (g *MphMatcherGroup) Match(input string) []uint32 { - matches := make([][]uint32, 0, 5) - hash := uint32(0) - for i := len(input) - 1; i >= 0; i-- { - hash = hash*PrimeRK + uint32(input[i]) - if input[i] == '.' { - if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 { - matches = append(matches, g.values[mphIdx]) +// appendValues appends the values of record e for the flags in want, in mphKinds order. +func (g *MphMatcherGroup) appendValues(dst []uint32, e, want uint32) []uint32 { + if !g.multi { + for _, flag := range mphKinds { + if e&want&flag != 0 { + dst = append(dst, g.single) + } + } + return dst + } + if e&want == 0 { + return dst + } + p, n := g.recSpan(e & mphOffMask) + p += n + for _, flag := range mphKinds { + if e&flag == 0 { + continue + } + var count, v uint32 + for count, p = g.uvarint(p); count > 0; count-- { + v, p = g.uvarint(p) + if want&flag != 0 { + dst = append(dst, v) } } } - if mphIdx := g.Lookup(hash, input); mphIdx != 0 { - matches = append(matches, g.values[mphIdx]) + return dst +} + +// 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. 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-- { - hash = hash*PrimeRK + uint32(input[i]) if input[i] == '.' { - if g.Lookup(hash, input[i:]) != 0 { - return true - } + dst = append(dst, mphSuffix{h, i + 1}) + } + 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 diff --git a/common/geodata/strmatcher/matchergroup_mph_internal_test.go b/common/geodata/strmatcher/matchergroup_mph_internal_test.go new file mode 100644 index 000000000..f6b320f60 --- /dev/null +++ b/common/geodata/strmatcher/matchergroup_mph_internal_test.go @@ -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) + } +} diff --git a/common/geodata/strmatcher/matchergroup_mph_test.go b/common/geodata/strmatcher/matchergroup_mph_test.go index f710c5b94..acb82a72a 100644 --- a/common/geodata/strmatcher/matchergroup_mph_test.go +++ b/common/geodata/strmatcher/matchergroup_mph_test.go @@ -1,7 +1,10 @@ package strmatcher_test import ( + "math/rand" "reflect" + "slices" + "strings" "testing" "github.com/xtls/xray-core/common" @@ -276,3 +279,142 @@ func TestEmptyMphMatcherGroup(t *testing.T) { 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) +} diff --git a/common/geodata/strmatcher/matchers.go b/common/geodata/strmatcher/matchers.go index a9df9d6f2..e0f752ed4 100644 --- a/common/geodata/strmatcher/matchers.go +++ b/common/geodata/strmatcher/matchers.go @@ -2,9 +2,12 @@ package strmatcher import ( "errors" + "math/bits" "regexp" + "regexp/syntax" "slices" "strings" + "unicode" "unicode/utf8" "golang.org/x/net/idna" @@ -73,7 +76,274 @@ func (m SubstrMatcher) Match(s string) bool { // RegexMatcher is an implementation of Matcher. 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< 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 { @@ -89,6 +359,14 @@ func (m *RegexMatcher) String() string { } func (m *RegexMatcher) Match(s string) bool { + if !m.mayMatch(s) { + return false + } + for _, l := range m.literals { + if !strings.Contains(s, l) { + return false + } + } return m.pattern.MatchString(s) } @@ -102,11 +380,7 @@ func (t Type) New(pattern string) (Matcher, error) { case Domain: return DomainMatcher(pattern), nil case Regex: // 1. regex matching is case-sensitive - regex, err := regexp.Compile(pattern) - if err != nil { - return nil, err - } - return &RegexMatcher{pattern: regex}, nil + return newRegexMatcher(pattern) default: return nil, errors.New("unknown matcher type") } @@ -135,11 +409,7 @@ func (t Type) NewDomainPattern(pattern string) (Matcher, error) { } return DomainMatcher(pattern), nil case Regex: // Regex's charset not in LDH subset - regex, err := regexp.Compile(pattern) - if err != nil { - return nil, err - } - return &RegexMatcher{pattern: regex}, nil + return newRegexMatcher(pattern) default: return nil, errors.New("unknown matcher type") } diff --git a/common/geodata/strmatcher/matchers_regex_test.go b/common/geodata/strmatcher/matchers_regex_test.go new file mode 100644 index 000000000..cb6a92a3a --- /dev/null +++ b/common/geodata/strmatcher/matchers_regex_test.go @@ -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:]) + } + } + }) +} diff --git a/common/geodata/strmatcher/valuematcher_mph.go b/common/geodata/strmatcher/valuematcher_mph.go index ced6f442f..f764dc4d4 100644 --- a/common/geodata/strmatcher/valuematcher_mph.go +++ b/common/geodata/strmatcher/valuematcher_mph.go @@ -46,7 +46,9 @@ func (g *MphValueMatcher) Add(matcher Matcher, value uint32) { func (g *MphValueMatcher) Build() error { if g.mph != nil { runtime.GC() // peak mem - g.mph.Build() + if err := g.mph.Build(); err != nil { + return err + } } runtime.GC() // peak mem if g.ac != nil { @@ -58,23 +60,17 @@ func (g *MphValueMatcher) Build() error { // Match implements ValueMatcher.Match. func (g *MphValueMatcher) Match(input string) []uint32 { - result := make([][]uint32, 0, 5) + var result []uint32 if g.mph != nil { - if matches := g.mph.Match(input); len(matches) > 0 { - result = append(result, matches) - } + result = g.mph.Match(input) // a new slice, returned without another copy } if g.ac != nil { - if matches := g.ac.Match(input); len(matches) > 0 { - result = append(result, matches) - } + result = append(result, g.ac.Match(input)...) } if g.regex != nil { - if matches := g.regex.Match(input); len(matches) > 0 { - result = append(result, matches) - } + result = append(result, g.regex.Match(input)...) } - return CompositeMatches(result) + return result } // MatchAny implements ValueMatcher.MatchAny. @@ -87,3 +83,62 @@ func (g *MphValueMatcher) MatchAny(input string) bool { } return g.regex != nil && g.regex.MatchAny(input) } + +func (g *MphValueMatcher) matchAnyHashed(input string, parents []mphSuffix, h, mul uint64) bool { + if g.mph != nil && g.mph.matchAnyHashed(input, parents, h, mul) { + return true + } + if g.ac != nil && g.ac.MatchAny(input) { + return true + } + return g.regex != nil && g.regex.MatchAny(input) +} + +// MphValueMatcherCombiner combines several built MphValueMatchers, each bound to one value, and matches an input +// against them as their MatchAny would, hashing the input once for all of them. +type MphValueMatcherCombiner struct { + matchers []*MphValueMatcher + values []uint32 +} + +// Add adds a built matcher that stands for value. +func (s *MphValueMatcherCombiner) Add(m *MphValueMatcher, value uint32) { + s.matchers = append(s.matchers, m) + s.values = append(s.values, value) +} + +// Match returns the values of the matchers that match input, in Add order. +func (s *MphValueMatcherCombiner) Match(input string) []uint32 { + if len(s.matchers) == 0 { + return nil + } + var stack [16]mphSuffix + mul := mphMultipliers[0] + parents, h := mphSuffixes(stack[:0], mul, input) + var result []uint32 + for i, m := range s.matchers { + if m.matchAnyHashed(input, parents, h, mul) { + result = append(result, s.values[i]) + } + } + return result +} + +// MatchAny returns true as soon as one matcher matches input. +func (s *MphValueMatcherCombiner) MatchAny(input string) bool { + switch len(s.matchers) { + case 0: + return false + case 1: + return s.matchers[0].MatchAny(input) // nothing to share, and it stops at the first matching suffix + } + var stack [16]mphSuffix + mul := mphMultipliers[0] + parents, h := mphSuffixes(stack[:0], mul, input) + for _, m := range s.matchers { + if m.matchAnyHashed(input, parents, h, mul) { + return true + } + } + return false +} diff --git a/common/log/log.go b/common/log/log.go index fbc2a509e..9759df273 100644 --- a/common/log/log.go +++ b/common/log/log.go @@ -1,7 +1,7 @@ package log // import "github.com/xtls/xray-core/common/log" import ( - "sync" + "sync/atomic" "github.com/xtls/xray-core/common/serial" ) @@ -29,36 +29,32 @@ func (m *GeneralMessage) String() string { // Record writes a message into log stream. 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. func RegisterHandler(handler Handler) { if handler == nil { panic("Log handler is nil") } - logHandler.Set(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 + logHandler.Store(&handler) } diff --git a/common/log/logger.go b/common/log/logger.go index 1b4eda864..36ba438af 100644 --- a/common/log/logger.go +++ b/common/log/logger.go @@ -68,6 +68,10 @@ func (l *serverityLogger) Handle(msg Message) { } } +func (l *serverityLogger) Severity() Severity { + return l.logLevel +} + func (l *generalLogger) run() { defer l.access.Signal() diff --git a/common/mux/client.go b/common/mux/client.go index 283803312..0463be5fc 100644 --- a/common/mux/client.go +++ b/common/mux/client.go @@ -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 { diff --git a/common/mux/frame.go b/common/mux/frame.go index f248fbdf6..f30195416 100644 --- a/common/mux/frame.go +++ b/common/mux/frame.go @@ -117,7 +117,7 @@ func (f *FrameMetadata) Unmarshal(reader io.Reader, readSourceAndLocal bool) err return err } if metaLen > 512 { - return errors.New("invalid metalen ", metaLen).AtError() + return errors.New("invalid metalen ", metaLen) } b := buf.New() diff --git a/common/mux/server.go b/common/mux/server.go index d1cdac113..c23a3789e 100644 --- a/common/mux/server.go +++ b/common/mux/server.go @@ -351,7 +351,7 @@ func (w *ServerWorker) handleFrame(ctx context.Context, reader *buf.BufferedRead err = w.handleStatusKeep(&meta, reader) default: status := meta.SessionStatus - return errors.New("unknown status: ", status).AtError() + return errors.New("unknown status: ", status) } if err != nil { diff --git a/common/net/packet.go b/common/net/packet.go new file mode 100644 index 000000000..10fbaa6e9 --- /dev/null +++ b/common/net/packet.go @@ -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 +} diff --git a/common/protocol/user.go b/common/protocol/user.go index 75e8e6541..ebdcbf862 100644 --- a/common/protocol/user.go +++ b/common/protocol/user.go @@ -7,7 +7,7 @@ import ( func (u *User) GetTypedAccount() (Account, error) { if u.GetAccount() == nil { - return nil, errors.New("Account is missing").AtWarning() + return nil, errors.New("Account is missing") } rawAccount, err := u.Account.GetInstance() diff --git a/common/singbridge/destination.go b/common/singbridge/destination.go deleted file mode 100644 index 217c4d082..000000000 --- a/common/singbridge/destination.go +++ /dev/null @@ -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 -} diff --git a/common/singbridge/dialer.go b/common/singbridge/dialer.go deleted file mode 100644 index 07c428813..000000000 --- a/common/singbridge/dialer.go +++ /dev/null @@ -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 -} diff --git a/common/singbridge/error.go b/common/singbridge/error.go deleted file mode 100644 index ac9e63517..000000000 --- a/common/singbridge/error.go +++ /dev/null @@ -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 -} diff --git a/common/singbridge/handler.go b/common/singbridge/handler.go deleted file mode 100644 index ee4b7c15f..000000000 --- a/common/singbridge/handler.go +++ /dev/null @@ -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()) -} diff --git a/common/singbridge/logger.go b/common/singbridge/logger.go deleted file mode 100644 index 16ff29cc3..000000000 --- a/common/singbridge/logger.go +++ /dev/null @@ -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) { -} diff --git a/common/singbridge/packet.go b/common/singbridge/packet.go deleted file mode 100644 index fde4bed3d..000000000 --- a/common/singbridge/packet.go +++ /dev/null @@ -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 -} diff --git a/common/singbridge/pipe.go b/common/singbridge/pipe.go deleted file mode 100644 index 94c5ee0c1..000000000 --- a/common/singbridge/pipe.go +++ /dev/null @@ -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 -} diff --git a/common/singbridge/reader.go b/common/singbridge/reader.go deleted file mode 100644 index 1ace1845f..000000000 --- a/common/singbridge/reader.go +++ /dev/null @@ -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 -} diff --git a/common/type.go b/common/type.go index 8ee8745c7..d3de166fe 100644 --- a/common/type.go +++ b/common/type.go @@ -16,7 +16,7 @@ var typeCreatorRegistry = make(map[reflect.Type]ConfigCreator) func RegisterConfig(config interface{}, configCreator ConfigCreator) error { configType := reflect.TypeOf(config) 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 return nil @@ -27,7 +27,7 @@ func CreateObject(ctx context.Context, config interface{}) (interface{}, error) configType := reflect.TypeOf(config) creator, found := typeCreatorRegistry[configType] 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) } diff --git a/core/config.go b/core/config.go index 3ab2a3765..112035782 100644 --- a/core/config.go +++ b/core/config.go @@ -125,7 +125,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) { } 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" { @@ -142,7 +142,7 @@ func LoadConfig(formatName string, input interface{}) (*Config, error) { if len(v) == 1 { return configLoaderByName["protobuf"].Loader(v) } 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 { return f.Loader(v) } 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) { diff --git a/core/core.go b/core/core.go index 45b7c4ed9..474c52c40 100644 --- a/core/core.go +++ b/core/core.go @@ -20,7 +20,7 @@ import ( var ( Version_x byte = 26 Version_y byte = 9 - Version_z byte = 9 + Version_z byte = 30 ) var ( diff --git a/features/dns/fakedns.go b/features/dns/fakedns.go index 7aff1fcb5..021f820b9 100644 --- a/features/dns/fakedns.go +++ b/features/dns/fakedns.go @@ -13,7 +13,7 @@ type FakeDNSEngine interface { var ( FakeIPv4Pool = "198.18.0.0/15" - FakeIPv6Pool = "fc00::/18" + FakeIPv6Pool = "2001:2::/48" ) type FakeDNSEngineRev0 interface { diff --git a/features/dns/localdns/client.go b/features/dns/localdns/client.go index b00febb32..41480b179 100644 --- a/features/dns/localdns/client.go +++ b/features/dns/localdns/client.go @@ -97,6 +97,9 @@ func New() *Client { r := &net.Resolver{ PreferGo: true, 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) }, } diff --git a/features/dns/localdns/client_test.go b/features/dns/localdns/client_test.go new file mode 100644 index 000000000..342140515 --- /dev/null +++ b/features/dns/localdns/client_test.go @@ -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() +} diff --git a/go.mod b/go.mod index 6b2e5a405..fd55450e8 100644 --- a/go.mod +++ b/go.mod @@ -18,8 +18,6 @@ require ( github.com/pires/go-proxyproto v0.15.0 github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af 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/vishvananda/netlink v1.3.1 github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 diff --git a/go.sum b/go.sum index 79eff3694..d802e7cbd 100644 --- a/go.sum +++ b/go.sum @@ -73,10 +73,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/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/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/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0= diff --git a/infra/conf/http.go b/infra/conf/http.go index 449b2bf2e..3e1482cb5 100644 --- a/infra/conf/http.go +++ b/infra/conf/http.go @@ -97,7 +97,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) { user.Email = v.Email } else { 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) @@ -106,7 +106,7 @@ func (v *HTTPClientConfig) Build() (proto.Message, error) { account.Password = v.Password } else { 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()) diff --git a/infra/conf/lint.go b/infra/conf/lint.go index f8a6b38c1..90e276eb8 100644 --- a/infra/conf/lint.go +++ b/infra/conf/lint.go @@ -18,7 +18,7 @@ func RegisterConfigureFilePostProcessingStage(name string, stage ConfigureFilePo func PostProcessConfigureFile(conf *Config) error { for k, v := range configureFilePostProcessingStages { 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 diff --git a/infra/conf/loader.go b/infra/conf/loader.go index 1dc2de23b..36e3e1c77 100644 --- a/infra/conf/loader.go +++ b/infra/conf/loader.go @@ -13,7 +13,7 @@ type ConfigCreatorCache map[string]ConfigCreator func (v ConfigCreatorCache) RegisterCreator(id string, creator ConfigCreator) error { if _, found := v[id]; found { - return errors.New(id, " already registered.").AtError() + return errors.New(id, " already registered.") } v[id] = creator @@ -61,7 +61,7 @@ func (v *JSONConfigLoader) Load(raw []byte) (interface{}, string, error) { } rawID, found := obj[v.idKey] 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 if err := json.Unmarshal(rawID, &id); err != nil { diff --git a/infra/conf/masque.go b/infra/conf/masque.go index cb9c201a2..b128ee593 100644 --- a/infra/conf/masque.go +++ b/infra/conf/masque.go @@ -2,9 +2,11 @@ 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" ) @@ -35,3 +37,67 @@ func (c *MasqueClientConfig) Build() (proto.Message, error) { 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 +} diff --git a/infra/conf/masque_test.go b/infra/conf/masque_test.go index 806bb4842..14ec6eedf 100644 --- a/infra/conf/masque_test.go +++ b/infra/conf/masque_test.go @@ -4,7 +4,10 @@ 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" ) @@ -37,6 +40,14 @@ func TestMasqueConfig(t *testing.T) { 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{ @@ -47,6 +58,8 @@ func TestMasqueConfig(t *testing.T) { `{"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) @@ -83,3 +96,93 @@ func TestMasqueOutboundConfig(t *testing.T) { } } } + +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") + } +} diff --git a/infra/conf/serial/builder.go b/infra/conf/serial/builder.go index 755b3469c..3154fce3c 100644 --- a/infra/conf/serial/builder.go +++ b/infra/conf/serial/builder.go @@ -30,7 +30,7 @@ func MergeConfigFromFiles(files []*core.ConfigSource) (string, error) { if j, ok := creflect.MarshalToJson(c, true); ok { 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) { diff --git a/infra/conf/shadowsocks.go b/infra/conf/shadowsocks.go index 18451ab5c..13862dd83 100644 --- a/infra/conf/shadowsocks.go +++ b/infra/conf/shadowsocks.go @@ -3,8 +3,6 @@ package conf import ( "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/protocol" "github.com/xtls/xray-core/common/serial" @@ -55,7 +53,7 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) { v.Users = v.Clients } - if C.Contains(shadowaead_2022.List, v.Cipher) { + if _, err := shadowsocks_2022.GetCipherMethod(v.Cipher); err == nil { return buildShadowsocks2022(v) } @@ -111,12 +109,14 @@ func (v *ShadowsocksServerConfig) Build() (proto.Message, error) { } func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) { + v.Cipher = strings.ToLower(v.Cipher) if len(v.Users) == 0 { config := new(shadowsocks_2022.ServerConfig) config.Method = v.Cipher config.Key = v.Password config.Network = v.NetworkList.Build() config.Email = v.Email + config.Level = int32(v.Level) return config, nil } @@ -171,6 +171,7 @@ func buildShadowsocks2022(v *ShadowsocksServerConfig) (proto.Message, error) { Email: user.Email, Address: user.Address.Build(), Port: uint32(user.Port), + Level: int32(user.Level), }) } 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`) } - if len(v.Servers) == 1 { - server := v.Servers[0] - if C.Contains(shadowaead_2022.List, server.Cipher) { - 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.") - } - - 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 - } + server := v.Servers[0] + 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.") } + 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) - for _, server := range v.Servers { - if C.Contains(shadowaead_2022.List, server.Cipher) { - 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 := &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 return config, nil } diff --git a/infra/conf/socks.go b/infra/conf/socks.go index 828584f18..33e290b7b 100644 --- a/infra/conf/socks.go +++ b/infra/conf/socks.go @@ -44,7 +44,6 @@ func (v *SocksServerConfig) Build() (proto.Message, error) { case AuthMethodUserPass: config.AuthType = socks.AuthType_PASSWORD default: - // errors.New("unknown socks auth method: ", v.AuthMethod, ". Default to noauth.").AtWarning().WriteToLog() config.AuthType = socks.AuthType_NO_AUTH } @@ -115,7 +114,7 @@ func (v *SocksClientConfig) Build() (proto.Message, error) { user.Email = v.Email } else { 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) @@ -124,7 +123,7 @@ func (v *SocksClientConfig) Build() (proto.Message, error) { account.Password = v.Password } else { 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()) diff --git a/infra/conf/transport_finalmask.go b/infra/conf/transport_finalmask.go index c46016584..69a69b5f2 100644 --- a/infra/conf/transport_finalmask.go +++ b/infra/conf/transport_finalmask.go @@ -1,6 +1,7 @@ package conf import ( + "context" "crypto/x509" "encoding/base64" "encoding/hex" @@ -14,6 +15,7 @@ import ( googleuuid "github.com/google/uuid" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/serial" "github.com/xtls/xray-core/transport/internet/finalmask/fragment" "github.com/xtls/xray-core/transport/internet/finalmask/header/custom" "github.com/xtls/xray-core/transport/internet/finalmask/mkcp/aes128gcm" @@ -81,7 +83,7 @@ var ( "noise": func() interface{} { return new(NoiseMask) }, "salamander": func() interface{} { return new(Salamander) }, "sudoku": func() interface{} { return new(Sudoku) }, - "xdns": func() interface{} { return new(Xdns) }, + "xdns": func() interface{} { return new(XDNS) }, "xicmp": func() interface{} { return new(Xicmp) }, "realm": func() interface{} { return new(Realm) }, "udphop": func() interface{} { return new(UDPHop) }, @@ -308,14 +310,27 @@ type NoiseMask struct { } func (c *NoiseMask) Build() (proto.Message, error) { + noiseSlice := make([]*noise.Item, 0, len(c.Noise)) for _, item := range c.Noise { if len(item.Packet) > 0 && item.Rand.To > 0 { return nil, errors.New("len(item.Packet) > 0 && item.Rand.To > 0") } - } - - noiseSlice := make([]*noise.Item, 0, len(c.Noise)) - for _, item := range c.Noise { + if strings.ToLower(item.Type) == "exp" { + var exp string + if err := json.Unmarshal(item.Packet, &exp); err != nil { + return nil, errors.New(`"packet" of noise "type": "exp" must be a string`).Base(err) + } + segments, err := parseNoiseExp(exp) + if err != nil { + return nil, err + } + noiseSlice = append(noiseSlice, &noise.Item{ + Segments: segments, + DelayMin: int64(item.Delay.From), + DelayMax: int64(item.Delay.To), + }) + continue + } if item.RandRange == nil { item.RandRange = &Int32Range{From: 0, To: 255} } @@ -344,6 +359,88 @@ func (c *NoiseMask) Build() (proto.Message, error) { }, nil } +var noiseExpPattern = regexp.MustCompile(`<\s*([a-z]+)(?:\s+([^>]*?))?\s*>`) + +func parseNoiseExp(exp string) ([]*noise.Segment, error) { + var segments []*noise.Segment + matches := noiseExpPattern.FindAllStringSubmatchIndex(exp, -1) + last := 0 + for _, m := range matches { + if strings.TrimSpace(exp[last:m[0]]) != "" { + return nil, errors.New("invalid noise exp near ", exp[last:m[0]]) + } + last = m[1] + key := exp[m[2]:m[3]] + arg := "" + if m[4] >= 0 { + arg = exp[m[4]:m[5]] + } + segment, err := buildNoiseSegment(key, arg) + if err != nil { + return nil, err + } + segments = append(segments, segment) + } + if strings.TrimSpace(exp[last:]) != "" { + return nil, errors.New("invalid noise exp near ", exp[last:]) + } + if len(segments) == 0 { + return nil, errors.New("empty noise exp: ", exp) + } + return segments, nil +} + +func buildNoiseSegment(key, arg string) (*noise.Segment, error) { + sizeSegment := func(kind noise.Segment_Kind) (*noise.Segment, error) { + if arg == "" { + return nil, errors.New("<", key, "> in noise exp needs a size") + } + lo, hi, err := ParseRangeString(arg) + if err != nil { + return nil, err + } + if lo < 0 || hi < lo || hi > 65535 { + return nil, errors.New("invalid size in noise exp: ", arg) + } + return &noise.Segment{Kind: kind, MinSize: int64(lo), MaxSize: int64(hi)}, nil + } + switch key { + case "b": + hexStr := strings.TrimPrefix(strings.TrimPrefix(strings.Join(strings.Fields(arg), ""), "0x"), "0X") + if len(hexStr) == 0 { + return nil, errors.New("empty bytes in noise exp") + } + raw, err := hex.DecodeString(hexStr) + if err != nil { + return nil, errors.New("invalid hex in noise exp: ", arg).Base(err) + } + return &noise.Segment{Kind: noise.Segment_BYTES, Bytes: raw}, nil + case "r": + return sizeSegment(noise.Segment_RANDOM) + case "rc": + return sizeSegment(noise.Segment_RANDOM_ASCII) + case "rd": + return sizeSegment(noise.Segment_RANDOM_DIGIT) + case "t": + if arg != "" { + return nil, errors.New(" in noise exp takes no argument") + } + return &noise.Segment{Kind: noise.Segment_TIMESTAMP}, nil + case "c": + if arg != "" { + return nil, errors.New(" in noise exp takes no argument") + } + return &noise.Segment{Kind: noise.Segment_COUNTER}, nil + case "n": + if arg != "" { + return nil, errors.New(" in noise exp takes no argument") + } + return &noise.Segment{Kind: noise.Segment_NONCE}, nil + default: + return nil, errors.New("unknown <", key, "> in noise exp") + } +} + type UDPItem struct { Rand int32 `json:"rand"` RandRange *Int32Range `json:"randRange"` @@ -694,32 +791,88 @@ func (c *Sudoku) Build() (proto.Message, error) { }, nil } -type Xdns struct { - Domain json.RawMessage `json:"domain"` - - Domains []string `json:"domains"` - Resolvers []string `json:"resolvers"` +type XDNSDomain struct { + Name string `json:"name"` + LenLimit int32 `json:"lenLimit"` + LabelLimit int32 `json:"labelLimit"` + Types []int32 `json:"types"` + Edns0 int32 `json:"edns0"` } -func (c *Xdns) Build() (proto.Message, error) { - if c.Domain != nil { - return nil, errors.PrintRemovedFeatureError("domain", "domains(server) & resolvers(client)") - } +type XDNSResolverTCP struct { + Addr string `json:"addr"` +} - if len(c.Domains) == 0 && len(c.Resolvers) == 0 { - return nil, errors.New("empty domains & empty resolvers") - } +func (c *XDNSResolverTCP) Build() (proto.Message, error) { + return &xdns.TCPResolverProto{Addr: c.Addr}, nil +} - for _, r := range c.Resolvers { - if !strings.Contains(r, "+udp://") { - return nil, errors.New("invalid resolver ", r) +type XDNSResolverUDP struct { + Addr string `json:"addr"` +} + +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, + }) } - - return &xdns.Config{ - Domains: c.Domains, - Resolvers: c.Resolvers, - }, nil + for i := range c.Resolvers { + config, err := xdnsLoader.LoadWithID(c.Resolvers[i].Settings, c.Resolvers[i].Type) + if err != nil { + return nil, err + } + 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 { @@ -942,12 +1095,19 @@ func (c *UDPHop) Build() (proto.Message, error) { } 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{ Local: local, Remote: remote, RemoteOnce: remoteOnce, - IntervalMin: int64(c.Interval.From), - IntervalMax: int64(c.Interval.To), + IntervalMin: int64(interval.From), + IntervalMax: int64(interval.To), RemoteIPs: remoteIPs, RemotePorts: c.RemotePorts.Build().Ports(), }, nil diff --git a/infra/conf/transport_finalmask_noise_test.go b/infra/conf/transport_finalmask_noise_test.go new file mode 100644 index 000000000..cca8a8666 --- /dev/null +++ b/infra/conf/transport_finalmask_noise_test.go @@ -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("") + 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("") + 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(" "); err != nil { + t.Errorf("surrounding whitespace should be allowed: %v", err) + } + cfg, err := buildNoiseExp("") + 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{ + "", + "", + "", + "", + "", + "", + "", + "", + "", + "", + "garbage", + " tail", + "", + } { + 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(""), 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": "", "delay": "1-3"}, + {"type": "EXP", "packet": ""}, + {"type": "str", "packet": ""}, + {"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.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]) + } +} diff --git a/infra/conf/transport_method.go b/infra/conf/transport_method.go index ba671d8b3..eadc0519d 100644 --- a/infra/conf/transport_method.go +++ b/infra/conf/transport_method.go @@ -1,7 +1,9 @@ package conf import ( + "encoding/base64" "encoding/json" + "maps" "math/big" "net/url" "sort" @@ -124,7 +126,7 @@ func (v *AuthenticatorRequest) Build() (*http.RequestConfig, error) { for _, key := range headerNames { value := v.Headers[key] 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{ Name: key, @@ -192,7 +194,7 @@ func (v *AuthenticatorResponse) Build() (*http.ResponseConfig, error) { for _, key := range headerNames { value := v.Headers[key] 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{ Name: key, @@ -242,11 +244,11 @@ func (c *TCPConfig) Build() (proto.Message, error) { if len(c.HeaderConfig) > 0 { headerConfig, _, err := tcpHeaderLoader.Load(c.HeaderConfig) 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() 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) } @@ -791,6 +793,8 @@ func (c *HysteriaConfig) Build() (proto.Message, error) { 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"` } @@ -819,12 +823,27 @@ func (c *MasqueConfig) Build() (proto.Message, error) { 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: c.Headers, + Headers: headers, }, nil } diff --git a/infra/conf/tun.go b/infra/conf/tun.go index 73e71a991..a358b035c 100644 --- a/infra/conf/tun.go +++ b/infra/conf/tun.go @@ -5,8 +5,12 @@ import ( "fmt" "math/big" "net" + "runtime" + "slices" "strconv" + "strings" + "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/proxy/tun" "google.golang.org/protobuf/proto" ) @@ -20,6 +24,8 @@ type TunConfig struct { UserLevel uint32 `json:"userLevel"` AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"` AutoOutboundsInterface *string `json:"autoOutboundsInterface"` + AutoSystemDnsToGateway bool `json:"autoSystemDnsToGateway"` + AutoSystemWfpBlockLeak []string `json:"autoSystemWfpBlockLeak"` } func (v *TunConfig) Build() (proto.Message, error) { @@ -31,6 +37,32 @@ func (v *TunConfig) Build() (proto.Message, error) { DNS: v.DNS, UserLevel: v.UserLevel, 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 { config.AutoOutboundsInterface = *v.AutoOutboundsInterface diff --git a/infra/conf/tun_test.go b/infra/conf/tun_test.go new file mode 100644 index 000000000..8e8693b74 --- /dev/null +++ b/infra/conf/tun_test.go @@ -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") + } +} diff --git a/infra/conf/xray.go b/infra/conf/xray.go index 46d3ae2e4..1ce1bd06d 100644 --- a/infra/conf/xray.go +++ b/infra/conf/xray.go @@ -33,6 +33,7 @@ var ( "trojan": func() interface{} { return new(TrojanServerConfig) }, "wireguard": func() interface{} { return &WireGuardConfig{IsClient: false} }, "hysteria": func() interface{} { return new(HysteriaServerConfig) }, + "masque": func() interface{} { return new(MasqueServerConfig) }, "tun": func() interface{} { return new(TunConfig) }, }, "protocol", "settings") @@ -205,6 +206,9 @@ func (c *InboundDetourConfig) Build() (*core.InboundHandlerConfig, error) { if err != nil { 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{ Tag: c.Tag, diff --git a/main/commands/all/api/inbound_user_add.go b/main/commands/all/api/inbound_user_add.go index 81afc0d29..792d625eb 100644 --- a/main/commands/all/api/inbound_user_add.go +++ b/main/commands/all/api/inbound_user_add.go @@ -12,6 +12,8 @@ import ( "github.com/xtls/xray-core/core" "github.com/xtls/xray-core/infra/conf" "github.com/xtls/xray-core/infra/conf/serial" + "github.com/xtls/xray-core/proxy/hysteria" + "github.com/xtls/xray-core/proxy/masque" "github.com/xtls/xray-core/proxy/shadowsocks" "github.com/xtls/xray-core/proxy/shadowsocks_2022" "github.com/xtls/xray-core/proxy/trojan" @@ -88,6 +90,10 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User { return ty.Users case *shadowsocks_2022.MultiUserServerConfig: return ty.Users + case *masque.ServerConfig: + return ty.Users + case *hysteria.ServerConfig: + return ty.Users default: fmt.Println("unsupported inbound type") } diff --git a/proxy/freedom/freedom.go b/proxy/freedom/freedom.go index 47a33b560..59973594a 100644 --- a/proxy/freedom/freedom.go +++ b/proxy/freedom/freedom.go @@ -467,7 +467,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr if statConn != nil { counter = statConn.ReadCounter } - if c, ok := iConn.(*internet.PacketConnWrapper); ok { + if c, ok := iConn.(*net.PacketConnWrapper); ok { isOverridden := false if UDPOverride.Address != nil || UDPOverride.Port != 0 { isOverridden = true @@ -487,7 +487,7 @@ func NewPacketReader(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr } type PacketReader struct { - *internet.PacketConnWrapper + *net.PacketConnWrapper stats.Counter Handler *Handler DefaultRule *FinalRule @@ -542,7 +542,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr if statConn != nil { counter = statConn.WriteCounter } - if c, ok := iConn.(*internet.PacketConnWrapper); ok { + if c, ok := iConn.(*net.PacketConnWrapper); ok { // If DialDest is a domain, it will be resolved in dialer // check this behavior and add it to map resolvedUDPAddr := utils.NewTypedSyncMap[string, net.Address]() @@ -563,7 +563,7 @@ func NewPacketWriter(conn net.Conn, h *Handler, defaultRule *FinalRule, UDPOverr } type PacketWriter struct { - *internet.PacketConnWrapper + *net.PacketConnWrapper stats.Counter *Handler DefaultRule *FinalRule diff --git a/proxy/http/server.go b/proxy/http/server.go index c164542e9..b89c67c0e 100644 --- a/proxy/http/server.go +++ b/proxy/http/server.go @@ -115,11 +115,7 @@ Start: request, err := http.ReadRequest(reader) if err != nil { - trace := errors.New("failed to read http request").Base(err) - if errors.Cause(err) != io.EOF && !isTimeout(errors.Cause(err)) { - trace.AtWarning() - } - return trace + return errors.New("failed to read http request").Base(err) } if len(s.config.Accounts) > 0 { @@ -147,7 +143,7 @@ Start: } dest, err := http_proto.ParseHost(host, defaultPort) 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{ From: conn.RemoteAddr(), @@ -262,7 +258,7 @@ func (s *Server) handlePlainHTTP(ctx context.Context, request *http.Request, wri requestWriter := buf.NewBufferedWriter(link.Writer) common.Must(requestWriter.SetBuffered(false)) 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 } @@ -299,7 +295,7 @@ func (s *Server) handlePlainHTTP(ctx context.Context, request *http.Request, wri response.Header.Set("Proxy-Connection", "close") } 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 } diff --git a/proxy/hysteria/client.go b/proxy/hysteria/client.go index 7021ab5a5..cfced4289 100644 --- a/proxy/hysteria/client.go +++ b/proxy/hysteria/client.go @@ -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) 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() errors.LogInfo(ctx, "tunneling request to ", target, " via ", target.Network, ":", c.server.Destination.NetAddr()) diff --git a/proxy/hysteria/server.go b/proxy/hysteria/server.go index 1e54b654b..8debcc182 100644 --- a/proxy/hysteria/server.go +++ b/proxy/hysteria/server.go @@ -40,11 +40,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) { for _, user := range config.Users { u, err := user.ToMemoryUser() 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 { - return nil, errors.New("failed to add user").Base(err).AtError() + return nil, errors.New("failed to add user").Base(err) } } diff --git a/proxy/loopback/loopback.go b/proxy/loopback/loopback.go index 185496474..bf830e2e7 100644 --- a/proxy/loopback/loopback.go +++ b/proxy/loopback/loopback.go @@ -56,7 +56,7 @@ func (l *Loopback) init(config *Config, dispatcherInstance routing.Dispatcher) e if config.Sniffing.GetEnabled() { request, err := proxyman.BuildSniffingRequest(config.Sniffing) 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 } diff --git a/proxy/masque/account.go b/proxy/masque/account.go new file mode 100644 index 000000000..9b505bcf1 --- /dev/null +++ b/proxy/masque/account.go @@ -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)) +} diff --git a/proxy/masque/account_test.go b/proxy/masque/account_test.go new file mode 100644 index 000000000..907faa9f6 --- /dev/null +++ b/proxy/masque/account_test.go @@ -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()) +} diff --git a/proxy/masque/client.go b/proxy/masque/client.go index 9700943f7..161542f6b 100644 --- a/proxy/masque/client.go +++ b/proxy/masque/client.go @@ -151,7 +151,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter } defer conn.Close() uc := &wireguard.UDPConnClient{ - PacketConn: conn.(*internet.PacketConnWrapper).PacketConn, + PacketConn: conn.(*net.PacketConnWrapper).PacketConn, Dest: conn.RemoteAddr().(*net.UDPAddr), } reader = uc diff --git a/proxy/masque/config.pb.go b/proxy/masque/config.pb.go index 2d5d8e69c..bcb2d7c1f 100644 --- a/proxy/masque/config.pb.go +++ b/proxy/masque/config.pb.go @@ -74,15 +74,125 @@ func (x *ClientConfig) GetRemoteDns() []string { 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\"k\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\tremoteDnsBU\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 ( @@ -97,18 +207,22 @@ func file_proxy_masque_config_proto_rawDescGZIP() []byte { return file_proxy_masque_config_proto_rawDescData } -var file_proxy_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) +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 - (*protocol.ServerEndpoint)(nil), // 1: xray.common.protocol.ServerEndpoint + (*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{ - 1, // 0: xray.proxy.masque.ClientConfig.server:type_name -> xray.common.protocol.ServerEndpoint - 1, // [1:1] is the sub-list for method output_type - 1, // [1:1] is the sub-list for method input_type - 1, // [1:1] is the sub-list for extension type_name - 1, // [1:1] is the sub-list for extension extendee - 0, // [0:1] is the sub-list for field type_name + 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() } @@ -122,7 +236,7 @@ func file_proxy_masque_config_proto_init() { GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)), NumEnums: 0, - NumMessages: 1, + NumMessages: 3, NumExtensions: 0, NumServices: 0, }, diff --git a/proxy/masque/config.proto b/proxy/masque/config.proto index 0b402908a..3af0879dd 100644 --- a/proxy/masque/config.proto +++ b/proxy/masque/config.proto @@ -7,8 +7,19 @@ 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; +} diff --git a/proxy/masque/pool.go b/proxy/masque/pool.go new file mode 100644 index 000000000..3f24229d0 --- /dev/null +++ b/proxy/masque/pool.go @@ -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) +} diff --git a/proxy/masque/pool_test.go b/proxy/masque/pool_test.go new file mode 100644 index 000000000..f01421ba7 --- /dev/null +++ b/proxy/masque/pool_test.go @@ -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) + } +} diff --git a/proxy/masque/server.go b/proxy/masque/server.go new file mode 100644 index 000000000..f9cc2ff37 --- /dev/null +++ b/proxy/masque/server.go @@ -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)) + })) +} diff --git a/proxy/masque/server_test.go b/proxy/masque/server_test.go new file mode 100644 index 000000000..9b526f764 --- /dev/null +++ b/proxy/masque/server_test.go @@ -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) + } +} diff --git a/proxy/proxy.go b/proxy/proxy.go index ab7d7abc5..5cb9d7a11 100644 --- a/proxy/proxy.go +++ b/proxy/proxy.go @@ -277,6 +277,7 @@ func (w *VisionReader) ReadMultiBuffer() (buf.MultiBuffer, error) { w.ob.CanSpliceCopy = 1 } } + SuppressOuterCloseNotify(w.conn) readerConn, readCounter, _ := UnwrapRawConn(w.conn) w.directReadCounter = readCounter w.Reader = buf.NewReader(readerConn) @@ -340,6 +341,7 @@ func (w *VisionWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { // w.ob.CanSpliceCopy = 1 // } } + SuppressOuterCloseNotify(w.conn) rawConn, _, writerCounter := UnwrapRawConn(w.conn) w.Writer = buf.NewWriter(rawConn) w.directWriteCounter = writerCounter @@ -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 func UnwrapRawConn(conn net.Conn) (net.Conn, stats.Counter, stats.Counter) { var readCounter, writerCounter stats.Counter diff --git a/proxy/shadowsocks/client.go b/proxy/shadowsocks/client.go index 29cb04566..02ec72579 100644 --- a/proxy/shadowsocks/client.go +++ b/proxy/shadowsocks/client.go @@ -71,7 +71,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter return 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()) @@ -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 { - 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 { diff --git a/proxy/shadowsocks/protocol.go b/proxy/shadowsocks/protocol.go index 88006cded..354983f65 100644 --- a/proxy/shadowsocks/protocol.go +++ b/proxy/shadowsocks/protocol.go @@ -98,7 +98,7 @@ func ReadTCPSession(validator *Validator, reader io.Reader) (*protocol.RequestHe iv := append([]byte(nil), buffer.BytesTo(ivLen)...) r, err = account.Cipher.NewDecryptionReader(account.Key, iv, reader) 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) 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() diff --git a/proxy/shadowsocks/server.go b/proxy/shadowsocks/server.go index 360ea38c8..7c965b9f6 100644 --- a/proxy/shadowsocks/server.go +++ b/proxy/shadowsocks/server.go @@ -34,11 +34,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) { for _, user := range config.Users { u, err := user.ToMemoryUser() 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 { - 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 { sessionPolicy := s.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).AtWarning() + return errors.New("unable to set read deadline").Base(err) } bufferedReader := buf.BufferedReader{Reader: buf.NewReader(conn)} diff --git a/proxy/shadowsocks_2022/cipher.go b/proxy/shadowsocks_2022/cipher.go new file mode 100644 index 000000000..311b9bca9 --- /dev/null +++ b/proxy/shadowsocks_2022/cipher.go @@ -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") +} diff --git a/proxy/shadowsocks_2022/config.go b/proxy/shadowsocks_2022/config.go index 9ddd2cf88..2fa28aa0c 100644 --- a/proxy/shadowsocks_2022/config.go +++ b/proxy/shadowsocks_2022/config.go @@ -1,6 +1,9 @@ package shadowsocks_2022 import ( + "bytes" + "encoding/base64" + "google.golang.org/protobuf/proto" "github.com/xtls/xray-core/common/protocol" @@ -8,26 +11,31 @@ import ( // MemoryAccount is an account type converted from Account. type MemoryAccount struct { - Key string + Key []byte } // AsAccount implements protocol.AsAccount. func (u *Account) AsAccount() (protocol.Account, error) { + keyStr := u.GetKey() + raw, err := base64.StdEncoding.DecodeString(keyStr) + if err != nil { + raw = []byte(keyStr) + } return &MemoryAccount{ - Key: u.GetKey(), + Key: raw, }, nil } // Equals implements protocol.Account.Equals(). func (a *MemoryAccount) Equals(another protocol.Account) bool { if account, ok := another.(*MemoryAccount); ok { - return a.Key == account.Key + return bytes.Equal(a.Key, account.Key) } return false } func (a *MemoryAccount) ToProto() proto.Message { return &Account{ - Key: a.Key, + Key: base64.StdEncoding.EncodeToString(a.Key), } } diff --git a/proxy/shadowsocks_2022/inbound.go b/proxy/shadowsocks_2022/inbound.go index edf9857c8..9883e00e8 100644 --- a/proxy/shadowsocks_2022/inbound.go +++ b/proxy/shadowsocks_2022/inbound.go @@ -4,23 +4,16 @@ import ( "context" "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/antireplay" "github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/log" "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/protocol" "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/core" + "github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/routing" "github.com/xtls/xray-core/transport/internet/stat" ) @@ -32,10 +25,13 @@ func init() { } type Inbound struct { - networks []net.Network - service shadowsocks.Service - email string - level int + networks []net.Network + method *CipherMethod + psk []byte + user *protocol.MemoryUser + saltFilter *antireplay.ReplayFilter[[32]byte] + udpCodec *UDPServerCodec + policyManager policy.Manager } 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, } } - inbound := &Inbound{ - networks: networks, - 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) + + method, err := GetCipherMethod(config.Method) 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 { @@ -70,114 +81,105 @@ func (i *Inbound) Process(ctx context.Context, network net.Network, connection s inbound := session.InboundFromContext(ctx) inbound.Name = "shadowsocks-2022" inbound.CanSpliceCopy = 3 - - var metadata M.Metadata - if inbound.Source.IsValid() { - metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr()) - } - - ctx = session.ContextWithDispatcher(ctx, dispatcher) + inbound.User = i.user if network == net.Network_TCP { - return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata)) - } else { - reader := buf.NewReader(connection) - pc := &natPacketConn{connection} - for { - mb, err := reader.ReadMultiBuffer() - if err != nil { - buf.ReleaseMulti(mb) - return singbridge.ReturnError(err) + return i.processTCP(ctx, connection, dispatcher) + } + return i.processUDP(ctx, connection, dispatcher) +} + +func (i *Inbound) processTCP(ctx context.Context, conn net.Conn, dispatcher routing.Dispatcher) error { + defer conn.Close() + + sessionPolicy := i.policyManager.ForLevel(0) + if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { + return errors.New("unable to set read deadline").Base(err) + } + + // 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() - buffer.Release() - err = i.service.NewPacket(ctx, pc, packet, metadata) - if err != nil { - packet.Release() - buf.ReleaseMulti(mb) - return err + + sessionItem := i.udpCodec.GetSession(decoded.SessionID) + if sessionItem.User == nil { + sessionItem.Lock() + if sessionItem.User == nil { + sessionItem.User = i.user } + sessionItem.Unlock() } + link, err := sessionItem.EnsureLink(ctx, conn, decoded.Destination, dispatcher, i.policyManager, func(dest net.Destination, payload []byte) ([]byte, error) { + return i.udpCodec.EncodeServerPacket(decoded.SessionID, dest, payload) + }) + if err != nil { + 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 -} diff --git a/proxy/shadowsocks_2022/inbound_multi.go b/proxy/shadowsocks_2022/inbound_multi.go index d6d68c09a..746bd27f0 100644 --- a/proxy/shadowsocks_2022/inbound_multi.go +++ b/proxy/shadowsocks_2022/inbound_multi.go @@ -2,30 +2,26 @@ package shadowsocks_2022 import ( "context" - "encoding/base64" + "crypto/cipher" + "encoding/binary" "strconv" "strings" "sync" + "sync/atomic" "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/antireplay" "github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/log" "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/protocol" "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/utils" "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/transport/internet/stat" ) @@ -38,9 +34,16 @@ func init() { type MultiUserInbound struct { sync.Mutex - networks []net.Network - users []*protocol.MemoryUser - service *shadowaead_2022.MultiService[int] + networks []net.Network + method *CipherMethod + 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) { @@ -51,138 +54,131 @@ func NewMultiServer(ctx context.Context, config *MultiUserServerConfig) (*MultiU 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 == "" { 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 { - 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{ - 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 + return i, nil } -// AddUser implements proxy.UserManager.AddUser(). +// AddUser implements proxy.UserManager.AddUser() func (i *MultiUserInbound) AddUser(ctx context.Context, u *protocol.MemoryUser) error { i.Lock() defer i.Unlock() + var emailKey string if u.Email != "" { - for idx := range i.users { - if i.users[idx].Email == u.Email { - return errors.New("User ", u.Email, " already exists.") - } + emailKey = strings.ToLower(u.Email) + if _, exists := i.usersByEmail.Load(emailKey); exists { + return errors.New("user ", u.Email, " already exists") } } - i.users = append(i.users, u) - // 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 }), - ) + memAcc, ok := u.Account.(*MemoryAccount) + if !ok { + return errors.New("missing or invalid user account") + } + + 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 } -// RemoveUser implements proxy.UserManager.RemoveUser(). +// RemoveUser implements proxy.UserManager.RemoveUser() func (i *MultiUserInbound) RemoveUser(ctx context.Context, email string) error { if email == "" { - return errors.New("Email must not be empty.") + return errors.New("email must not be empty") } i.Lock() defer i.Unlock() - idx := -1 - for ii, u := range i.users { - if strings.EqualFold(u.Email, email) { - idx = ii - break - } + emailKey := strings.ToLower(email) + u, loaded := i.usersByEmail.LoadAndDelete(emailKey) + if !loaded { + return errors.New("user ", email, " not found") } - if idx == -1 { - return errors.New("User ", email, " not found.") - } - - 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 }), - ) + pskHash := DeriveUserPSKHash(u.Account.(*MemoryAccount).Key) + i.usersByHash.Delete(pskHash) + i.userCount.Add(-1) return nil } -// GetUser implements proxy.UserManager.GetUser(). +// GetUser implements proxy.UserManager.GetUser() func (i *MultiUserInbound) GetUser(ctx context.Context, email string) *protocol.MemoryUser { if email == "" { return nil } - - i.Lock() - defer i.Unlock() - - for _, u := range i.users { - if strings.EqualFold(u.Email, email) { - return u - } - } - return nil + u, _ := i.usersByEmail.Load(strings.ToLower(email)) + return u } -// GetUsers implements proxy.UserManager.GetUsers(). +// GetUsers implements proxy.UserManager.GetUsers() func (i *MultiUserInbound) GetUsers(ctx context.Context) []*protocol.MemoryUser { - i.Lock() - defer i.Unlock() - dst := make([]*protocol.MemoryUser, len(i.users)) - copy(dst, i.users) - return dst + var users []*protocol.MemoryUser + i.usersByEmail.Range(func(_ string, user *protocol.MemoryUser) bool { + users = append(users, user) + return true + }) + return users } -// GetUsersCount implements proxy.UserManager.GetUsersCount(). +// GetUsersCount implements proxy.UserManager.GetUsersCount() func (i *MultiUserInbound) GetUsersCount(context.Context) int64 { - i.Lock() - defer i.Unlock() - return int64(len(i.users)) + return i.userCount.Load() } 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.CanSpliceCopy = 3 - var metadata M.Metadata - if inbound.Source.IsValid() { - metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr()) + if network == net.Network_TCP { + return i.processTCP(ctx, connection, dispatcher) + } + 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 { - return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata)) - } else { - reader := buf.NewReader(connection) - pc := &natPacketConn{connection} - for { - mb, err := reader.ReadMultiBuffer() - if err != nil { - buf.ReleaseMulti(mb) - return singbridge.ReturnError(err) + var salt [32]byte + copy(salt[:i.method.KeySaltLength], headerBuf[:i.method.KeySaltLength]) + saltSlice := salt[:i.method.KeySaltLength] + eih := headerBuf[i.method.KeySaltLength : i.method.KeySaltLength+AESBlockSize] + fixedChunk := headerBuf[i.method.KeySaltLength+AESBlockSize:] + + decryptedHash, err := DecryptEIH(i.method, i.masterPSK, saltSlice, eih) + if err != nil { + ResetTCPConn(conn) + 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() - buffer.Release() - err = i.service.NewPacket(ctx, pc, packet, metadata) - if err != nil { - packet.Release() - buf.ReleaseMulti(mb) - return err + + var rawHeader [16]byte + i.udpMasterCipher.Decrypt(rawHeader[:], packetBytes[:16]) + + sessionID := binary.BigEndian.Uint64(rawHeader[:8]) + packetID := binary.BigEndian.Uint64(rawHeader[8:16]) + + 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 { - 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 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()) +func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) { + return i.udpSessions.EncodeServerPacket(i.method, userPSK, clientSessionID, dest, payload) } diff --git a/proxy/shadowsocks_2022/inbound_relay.go b/proxy/shadowsocks_2022/inbound_relay.go index 4ca5e2075..914cb91f0 100644 --- a/proxy/shadowsocks_2022/inbound_relay.go +++ b/proxy/shadowsocks_2022/inbound_relay.go @@ -2,18 +2,11 @@ package shadowsocks_2022 import ( "context" + "crypto/cipher" + "encoding/binary" "strconv" - "strings" "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/buf" "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/protocol" "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/core" + "github.com/xtls/xray-core/features/policy" "github.com/xtls/xray-core/features/routing" "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 { - networks []net.Network - destinations []*RelayDestination - service *shadowaead_2022.RelayService[int] + networks []net.Network + method *CipherMethod + relayPSK []byte + relayBlock cipher.Block + destinations map[[AESBlockSize]byte]*relayDest + udpSessions *UDPSessionManager + policyManager policy.Manager } func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) { @@ -48,39 +52,62 @@ func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbou net.Network_UDP, } } - inbound := &RelayInbound{ - networks: networks, - 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) + + method, err := GetCipherMethod(config.Method) 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 { - if destination.Email == "" { + relayPSK, err := ParseKey(config.Key, method.KeySaltLength) + 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() - 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 }), - 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 + + return i, nil } 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.CanSpliceCopy = 3 - var metadata M.Metadata - if inbound.Source.IsValid() { - metadata.Source = M.ParseSocksaddr(inbound.Source.NetAddr()) + if network == net.Network_TCP { + return i.processTCP(ctx, connection, dispatcher) + } + 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 { - return singbridge.ReturnError(i.service.NewConnection(ctx, connection, metadata)) - } else { - reader := buf.NewReader(connection) - pc := &natPacketConn{connection} - for { - mb, err := reader.ReadMultiBuffer() - if err != nil { - buf.ReleaseMulti(mb) - return singbridge.ReturnError(err) + headerSlice := requestHeader.Bytes() + salt := headerSlice[:i.method.KeySaltLength] + eih := headerSlice[i.method.KeySaltLength:needed] + + decryptedHash, err := DecryptEIH(i.method, i.relayPSK, salt, eih) + if err != nil { + requestHeader.Release() + ResetTCPConn(conn) + return 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() - buffer.Release() - err = i.service.NewPacket(ctx, pc, packet, metadata) - if err != nil { - packet.Release() - buf.ReleaseMulti(mb) - return err + + var packetHeader [AESBlockSize]byte + i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize]) + + eiHeader := DecryptUDPEIH(i.relayBlock, packetHeader[:], data[AESBlockSize:2*AESBlockSize]) + + targetDest, ok := i.destinations[eiHeader] + if !ok { + b.Release() + continue + } + + // Extract sessionID from raw packetHeader for session-level link caching before re-encrypting + sessionID := binary.BigEndian.Uint64(packetHeader[:8]) + + // Re-encrypt packetHeader with next hop block cipher + targetDest.blockCipher.Encrypt(packetHeader[:], packetHeader[:]) + + // Strip outer EIH: replace second block with re-encrypted packetHeader and advance + copy(data[AESBlockSize:2*AESBlockSize], packetHeader[:]) + b.Advance(int32(AESBlockSize)) + + dest := targetDest.destination + dest.Network = net.Network_UDP + + 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()) -} diff --git a/proxy/shadowsocks_2022/kdf.go b/proxy/shadowsocks_2022/kdf.go new file mode 100644 index 000000000..0f70c49df --- /dev/null +++ b/proxy/shadowsocks_2022/kdf.go @@ -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 +} diff --git a/proxy/shadowsocks_2022/outbound.go b/proxy/shadowsocks_2022/outbound.go index 5d1b9c9fb..67e0acd28 100644 --- a/proxy/shadowsocks_2022/outbound.go +++ b/proxy/shadowsocks_2022/outbound.go @@ -2,21 +2,19 @@ package shadowsocks_2022 import ( "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/buf" "github.com/xtls/xray-core/common/errors" "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/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/internet" ) @@ -28,42 +26,51 @@ func init() { } type Outbound struct { - ctx context.Context - server net.Destination - method shadowsocks.Method + server net.Destination + method *CipherMethod + pskList [][]byte + finalPSK []byte + udpCodec *UDPPacketCodec + policyManager policy.Manager } func NewClient(ctx context.Context, config *ClientConfig) (*Outbound, error) { - o := &Outbound{ - ctx: ctx, + method, err := GetCipherMethod(config.Method) + 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{ Address: config.Address.AsAddress(), Port: net.Port(config.Port), Network: net.Network_TCP, }, - } - if C.Contains(shadowaead_2022.List, config.Method) { - if config.Key == "" { - return nil, errors.New("missing psk") - } - method, err := shadowaead_2022.NewWithPassword(config.Method, config.Key, 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 + method: method, + pskList: pskList, + finalPSK: finalPSK, + udpCodec: udpCodec, + policyManager: v.GetFeature(policy.ManagerType()).(policy.Manager), + }, nil } 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) ob := outbounds[len(outbounds)-1] if !ob.Target.IsValid() { @@ -78,70 +85,140 @@ func (o *Outbound) Process(ctx context.Context, link *transport.Link, dialer int serverDestination := o.server 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) { - 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 { - serverConn := o.method.DialEarlyConn(connection, singbridge.ToSocksaddr(destination)) - var handshake bool - if timeoutReader, isTimeoutReader := link.Reader.(buf.TimeoutReader); isTimeoutReader { - mb, err := timeoutReader.ReadMultiBufferTimeout(time.Millisecond * 100) - 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), - } + var clientSalt [32]byte + clientSaltSlice := clientSalt[:o.method.KeySaltLength] + if _, err := io.ReadFull(rand.Reader, clientSaltSlice); err != nil { + return errors.New("failed to generate client salt").Base(err) } - serverConn := o.method.DialPacketConn(connection) - return singbridge.ReturnError(bufio.CopyPacketConn(ctx, packetConn, serverConn)) + requestDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) + + var initialPayload []byte + var firstBuf *buf.Buffer + var remainingMB buf.MultiBuffer + if timeoutReader, ok := link.Reader.(buf.TimeoutReader); ok { + if mb, err := timeoutReader.ReadMultiBufferTimeout(0); err == nil && !mb.IsEmpty() { + remainingMB, firstBuf = buf.SplitFirst(mb) + initialPayload = firstBuf.Bytes() + } + } + + bodyWriter, err := WriteTCPRequest(conn, o.method, o.pskList, destination, clientSaltSlice, initialPayload) + if firstBuf != nil { + firstBuf.Release() + } + 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) } diff --git a/proxy/shadowsocks_2022/packet.go b/proxy/shadowsocks_2022/packet.go new file mode 100644 index 000000000..f81754575 --- /dev/null +++ b/proxy/shadowsocks_2022/packet.go @@ -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 + } +} diff --git a/proxy/shadowsocks_2022/relay_test.go b/proxy/shadowsocks_2022/relay_test.go new file mode 100644 index 000000000..096852ca6 --- /dev/null +++ b/proxy/shadowsocks_2022/relay_test.go @@ -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)) + } + }) + } +} diff --git a/proxy/shadowsocks_2022/replay.go b/proxy/shadowsocks_2022/replay.go new file mode 100644 index 000000000..241eb053c --- /dev/null +++ b/proxy/shadowsocks_2022/replay.go @@ -0,0 +1,183 @@ +package shadowsocks_2022 + +import ( + "crypto/cipher" + "sync" + "sync/atomic" + "time" + + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/signal" + "github.com/xtls/xray-core/common/utils" + "github.com/xtls/xray-core/transport" +) + +const ( + swBlockBitLog = 6 // 1<<6 == 64 bits + swBlockBits = 1 << swBlockBitLog // 64 + swRingBlocks = 1 << 7 // 128 + swBlockMask = swRingBlocks - 1 // 127 + swBitMask = swBlockBits - 1 // 63 + swSize = (swRingBlocks - 1) * swBlockBits // 8128 +) + +type SlidingWindow struct { + last uint64 + ring [swRingBlocks]uint64 +} + +func (f *SlidingWindow) Reset() { + *f = SlidingWindow{} +} + +func (f *SlidingWindow) Check(counter uint64) bool { + switch { + case counter > f.last: + return true + case f.last-counter > swSize: + return false + } + + blockIndex := (counter >> swBlockBitLog) & swBlockMask + bitIndex := counter & swBitMask + return (f.ring[blockIndex]>>bitIndex)&1 == 0 +} + +func (f *SlidingWindow) Add(counter uint64) { + blockIndex := counter >> swBlockBitLog + + if counter > f.last { + lastBlockIndex := f.last >> swBlockBitLog + diff := int(blockIndex - lastBlockIndex) + if diff > swRingBlocks { + diff = swRingBlocks + } + + for i := 0; i < diff; i++ { + lastBlockIndex = (lastBlockIndex + 1) & swBlockMask + f.ring[lastBlockIndex] = 0 + } + + f.last = counter + } + + blockIndex &= swBlockMask + bitIndex := counter & swBitMask + f.ring[blockIndex] |= 1 << bitIndex +} + +func (f *SlidingWindow) CheckAndAdd(counter uint64) bool { + if !f.Check(counter) { + return false + } + f.Add(counter) + return true +} + +type ServerUDPSession struct { + sync.Mutex + SessionID uint64 + Window *SlidingWindow + User *protocol.MemoryUser + UserPSK []byte + LastActive atomic.Int64 // Unix timestamp in seconds + + clientBodyCipher cipher.AEAD + + ServerSessionID uint64 + ServerPacketID atomic.Uint64 + serverBodyCipher cipher.AEAD + serverHeaderBlock cipher.Block + serverChaCha cipher.AEAD + + manager *UDPSessionManager + link atomic.Pointer[transport.Link] + timer *signal.ActivityTimer + currentConn atomic.Value // stores stat.Connection +} + +func (s *ServerUDPSession) CheckPacketID(packetID uint64) bool { + s.Lock() + defer s.Unlock() + if s.Window == nil { + s.Window = new(SlidingWindow) + } + return s.Window.Check(packetID) +} + +func (s *ServerUDPSession) AddPacketID(packetID uint64) { + s.Lock() + defer s.Unlock() + if s.Window == nil { + s.Window = new(SlidingWindow) + } + s.Window.Add(packetID) +} + +type UDPSessionManager struct { + sessions *utils.TypedSyncMap[uint64, *ServerUDPSession] + timeout time.Duration + lastClean atomic.Int64 // Unix timestamp in seconds +} + +func NewUDPSessionManager(timeout time.Duration) *UDPSessionManager { + return &UDPSessionManager{ + sessions: utils.NewTypedSyncMap[uint64, *ServerUDPSession](), + timeout: timeout, + } +} + +func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession { + now := time.Now().Unix() + if s, ok := m.sessions.Load(sessionID); ok { + s.LastActive.Store(now) + return s + } + + s := &ServerUDPSession{ + SessionID: sessionID, + manager: m, + } + s.LastActive.Store(now) + + actual, loaded := m.sessions.LoadOrStore(sessionID, s) + if loaded { + actual.LastActive.Store(now) + return actual + } + + // Trigger cleanup if at least 30 seconds have passed since last cleanup + last := m.lastClean.Load() + if now-last > 30 && m.lastClean.CompareAndSwap(last, now) { + go m.cleanup(now) + } + + return s +} + +func (m *UDPSessionManager) cleanup(now int64) { + timeoutSec := int64(m.timeout.Seconds()) + if timeoutSec <= 0 { + timeoutSec = 60 + } + m.sessions.Range(func(k uint64, v *ServerUDPSession) bool { + if now-v.LastActive.Load() > timeoutSec { + m.sessions.Delete(k) + v.Close() + } + return true + }) +} + +func (m *UDPSessionManager) Delete(sessionID uint64) { + m.sessions.Delete(sessionID) +} + +func (m *UDPSessionManager) EncodeServerPacket(method *CipherMethod, psk []byte, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) { + sessionItem := m.GetOrCreate(clientSessionID) + if err := sessionItem.EnsureServerState(method, psk); err != nil { + return nil, err + } + return sessionItem.EncodeServerPacket(method, clientSessionID, dest, payload) +} diff --git a/proxy/shadowsocks_2022/shadowsocks_2022.go b/proxy/shadowsocks_2022/shadowsocks_2022.go index 96f62c74b..ec8562e6a 100644 --- a/proxy/shadowsocks_2022/shadowsocks_2022.go +++ b/proxy/shadowsocks_2022/shadowsocks_2022.go @@ -1 +1,193 @@ package shadowsocks_2022 + +import ( + "context" + + "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/buf" + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/log" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/session" + "github.com/xtls/xray-core/common/signal" + "github.com/xtls/xray-core/features/policy" + "github.com/xtls/xray-core/features/routing" + "github.com/xtls/xray-core/proxy" + "github.com/xtls/xray-core/transport" + "github.com/xtls/xray-core/transport/internet/stat" +) + +func (s *ServerUDPSession) UpdateConn(conn stat.Connection) { + if s.currentConn.Load() == nil { + s.currentConn.Store(conn) + } + if s.timer != nil { + s.timer.Update() + } +} + +func (s *ServerUDPSession) WriteToClient(b []byte) error { + connVal := s.currentConn.Load() + if connVal == nil { + return errors.New("client connection closed") + } + conn, ok := connVal.(stat.Connection) + if !ok || conn == nil { + return errors.New("client connection closed") + } + _, err := conn.Write(b) + return err +} + +func (s *ServerUDPSession) Close() { + if s.timer != nil { + s.timer.SetTimeout(0) + } + if link := s.link.Load(); link != nil { + common.Interrupt(link.Reader) + common.Interrupt(link.Writer) + } +} + +func (s *ServerUDPSession) EnsureLink( + ctx context.Context, + conn stat.Connection, + dest net.Destination, + dispatcher routing.Dispatcher, + policyManager policy.Manager, + responseEncoder func(dest net.Destination, payload []byte) ([]byte, error), +) (*transport.Link, error) { + s.UpdateConn(conn) + + if link := s.link.Load(); link != nil { + return link, nil + } + + s.Lock() + defer s.Unlock() + + if link := s.link.Load(); link != nil { + return link, nil + } + + sessCtx, cancel := context.WithCancel(ctx) + inbound := session.InboundFromContext(sessCtx) + if inbound != nil && s.User != nil { + inbound.User = s.User + } + var email string + var level uint32 + if s.User != nil { + email = s.User.Email + level = s.User.Level + } + sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{ + From: conn.RemoteAddr(), + To: dest, + Status: log.AccessAccepted, + Email: email, + }) + + link, err := dispatcher.Dispatch(sessCtx, dest) + if err != nil { + cancel() + return nil, err + } + + s.link.Store(link) + sessionPolicy := policyManager.ForLevel(level) + s.timer = signal.CancelAfterInactivity(sessCtx, func() { + if s.manager != nil { + s.manager.Delete(s.SessionID) + } + s.Close() + cancel() + }, sessionPolicy.Timeouts.ConnectionIdle) + + go handleUDPResponse(s, link, dest, responseEncoder) + return link, nil +} + +// ResetTCPConn sets SO_LINGER to 0 per SIP022 §3.1.4 to consistently send RST on close +// when handshake or header validation fails. +func ResetTCPConn(conn net.Conn) { + rawConn, _, _ := proxy.UnwrapRawConn(conn) + if tcpConn, ok := rawConn.(*net.TCPConn); ok { + _ = tcpConn.SetLinger(0) + } +} + +func handleUDPResponse(s *ServerUDPSession, link *transport.Link, fallbackDest net.Destination, encode func(dest net.Destination, payload []byte) ([]byte, error)) { + defer func() { + if s.timer != nil { + s.timer.SetTimeout(0) + } + }() + for { + resMb, err := link.Reader.ReadMultiBuffer() + if err != nil { + return + } + if s.timer != nil { + s.timer.Update() + } + for i, rb := range resMb { + b := rb.Bytes() + if encode != nil { + replyDest := fallbackDest + if rb.UDP != nil { + replyDest = *rb.UDP + } + encPacket, err := encode(replyDest, b) + rb.Release() + if err != nil { + continue + } + if err := s.WriteToClient(encPacket); err != nil { + buf.ReleaseMulti(resMb[i+1:]) + return + } + } else { + err := s.WriteToClient(b) + rb.Release() + if err != nil { + buf.ReleaseMulti(resMb[i+1:]) + return + } + } + } + } +} + +const ( + HeaderTypeClient = 0 + HeaderTypeServer = 1 + MaxPaddingLength = 900 + PacketNonceSize = 24 + MaxPacketSize = 65535 + RequestHeaderFixedChunkLength = 1 + 8 + 2 // Type (1B) + Timestamp (8B) + VarHeaderLen (2B) + PacketMinimalHeaderSize = 30 + StreamNonceSize = 12 + AESBlockSize = 16 + AEADTagSize = 16 +) + +var zeroPadding [MaxPaddingLength]byte + +const ( + MethodAES128GCM = "2022-blake3-aes-128-gcm" + MethodAES256GCM = "2022-blake3-aes-256-gcm" + MethodChaCha20Poly1305 = "2022-blake3-chacha20-poly1305" +) + +var ( + ErrBadKey = errors.New("bad key") + ErrBadHeaderType = errors.New("bad header type") + ErrBadTimestamp = errors.New("bad timestamp") + ErrSaltNotUnique = errors.New("salt not unique") + ErrPacketIdNotUnique = errors.New("packet id not unique") + ErrPacketTooShort = errors.New("packet too short") + ErrPacketTooLarge = errors.New("packet too large") + ErrNoPadding = errors.New("bad request: missing payload or padding") + ErrInvalidRequest = errors.New("invalid request") +) diff --git a/proxy/shadowsocks_2022/shadowsocks_2022_test.go b/proxy/shadowsocks_2022/shadowsocks_2022_test.go new file mode 100644 index 000000000..382696efb --- /dev/null +++ b/proxy/shadowsocks_2022/shadowsocks_2022_test.go @@ -0,0 +1,498 @@ +package shadowsocks_2022_test + +import ( + "bytes" + "context" + "crypto/rand" + "encoding/base64" + "encoding/binary" + "io" + gonet "net" + "sync" + "testing" + "time" + + "github.com/google/go-cmp/cmp" + "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/buf" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/serial" + "github.com/xtls/xray-core/common/session" + "github.com/xtls/xray-core/core" + . "github.com/xtls/xray-core/proxy/shadowsocks_2022" +) + +func newTestContext() context.Context { + v, err := core.New(&core.Config{}) + common.Must(err) + ctx := context.WithValue(context.Background(), core.XrayKey(1), v) + ctx = session.ContextWithInbound(ctx, &session.Inbound{}) + return ctx +} + +func generateRandomKey(size int) string { + b := make([]byte, size) + _, _ = rand.Read(b) + return base64.StdEncoding.EncodeToString(b) +} + +func TestKDF(t *testing.T) { + // Test ParseKey + if _, err := ParseKey("", 16); err != ErrBadKey { + t.Fatalf("expected ErrBadKey for empty key, got %v", err) + } + + shortKey := base64.StdEncoding.EncodeToString([]byte("short")) + if _, err := ParseKey(shortKey, 16); err != ErrBadKey { + t.Fatalf("expected ErrBadKey for short key, got %v", err) + } + + exactKey := []byte("0123456789abcdef") + exactKeyB64 := base64.StdEncoding.EncodeToString(exactKey) + normExact, err := ParseKey(exactKeyB64, 16) + if err != nil || !bytes.Equal(normExact, exactKey) { + t.Fatalf("unexpected parsed exact key: %v, err: %v", normExact, err) + } + + longKey := base64.StdEncoding.EncodeToString([]byte("0123456789abcdef_longer_key_for_testing")) + if _, err := ParseKey(longKey, 16); err != ErrBadKey { + t.Fatalf("expected ErrBadKey for long key, got %v", err) + } + + // Test Session Subkey determinism + salt := []byte("random_salt_1234") + k1 := DeriveSessionSubKey(normExact, salt, 16) + k2 := DeriveSessionSubKey(normExact, salt, 16) + if !bytes.Equal(k1, k2) { + t.Fatal("DeriveSessionSubKey should be deterministic") + } + + // Identity subkey must differ from session subkey with same inputs + idKey := DeriveIdentitySubKey(normExact, salt, 16) + if bytes.Equal(k1, idKey) { + t.Fatal("DeriveIdentitySubKey must differ from DeriveSessionSubKey") + } + + // User PSK hash + h1 := DeriveUserPSKHash(normExact) + h2 := DeriveUserPSKHash(normExact) + if h1 != h2 { + t.Fatal("DeriveUserPSKHash should be deterministic") + } +} + +func TestSlidingWindow(t *testing.T) { + var window SlidingWindow + if !window.Check(1) { + t.Fatal("packet 1 should be accepted") + } + window.Add(1) + + if window.Check(1) { + t.Fatal("duplicate packet 1 should be rejected") + } + + if !window.Check(100) { + t.Fatal("packet 100 should be accepted") + } + window.Add(100) + + if window.Check(100) { + t.Fatal("duplicate packet 100 should be rejected") + } + + if !window.Check(50) { + t.Fatal("out-of-order packet 50 within window should be accepted") + } + window.Add(50) + if window.Check(50) { + t.Fatal("duplicate packet 50 should be rejected") + } + + // Check packet far behind window (> 8128) + window.Add(10000) + if window.Check(1) { + t.Fatal("packet 1 should be rejected as behind window") + } +} + +func TestTCPStream(t *testing.T) { + methods := []struct { + name string + keySize int + }{ + {MethodAES128GCM, 16}, + {MethodAES256GCM, 32}, + {MethodChaCha20Poly1305, 32}, + } + + dest := net.TCPDestination(net.LocalHostIP, net.Port(8080)) + testPayload := []byte("Hello, Shadowsocks 2022 Native Implementation!") + + for _, m := range methods { + t.Run(m.name, func(t *testing.T) { + rawKey := make([]byte, m.keySize) + _, _ = rand.Read(rawKey) + method, err := GetCipherMethod(m.name) + common.Must(err) + + clientConn, serverConn := gonet.Pipe() + defer clientConn.Close() + defer serverConn.Close() + + var wg sync.WaitGroup + wg.Add(2) + + var receivedDest net.Destination + var receivedPayload []byte + + // Server goroutine + go func() { + defer wg.Done() + salt := make([]byte, method.KeySaltLength) + _, err := io.ReadFull(serverConn, salt) + common.Must(err) + + sessionKey := DeriveSessionSubKey(rawKey, salt, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + common.Must(err) + + reader := NewStreamReader(serverConn, aead) + + // Read fixed chunk (11 + 16 bytes) + var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte + _, err = io.ReadFull(serverConn, fixedBuf[:]) + common.Must(err) + + plainFixed, err := aead.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil) + common.Must(err) + IncreaseNonce(reader.Nonce()) + if plainFixed[0] != HeaderTypeClient { + t.Errorf("expected client header type, got %d", plainFixed[0]) + } + + // Read variable chunk + varLen := int(plainFixed[9])<<8 | int(plainFixed[10]) + varBuf := make([]byte, varLen+AEADTagSize) + _, err = io.ReadFull(serverConn, varBuf) + common.Must(err) + + plainVar, err := aead.Open(varBuf[:0], reader.Nonce(), varBuf, nil) + common.Must(err) + IncreaseNonce(reader.Nonce()) + + dest, addrLen, err := ParseAddressPort(plainVar) + common.Must(err) + receivedDest = net.TCPDestination(dest.Address, dest.Port) + plainVar = plainVar[addrLen:] + padLen := int(binary.BigEndian.Uint16(plainVar[:2])) + receivedPayload = plainVar[2+padLen:] + + // Server sends response stream with receivedPayload as first payload + writer := NewServerStreamWriter(serverConn, method, rawKey, salt) + pBuf := buf.New() + pBuf.Write(receivedPayload) + _ = writer.WriteMultiBuffer(buf.MultiBuffer{pBuf}) + + // Read and echo additional stream data + mb, err := reader.ReadMultiBuffer() + common.Must(err) + _ = writer.WriteMultiBuffer(mb) + _ = writer.Close() + }() + + // Client goroutine + go func() { + defer wg.Done() + clientSalt := make([]byte, method.KeySaltLength) + common.Must2(io.ReadFull(rand.Reader, clientSalt)) + writer, err := WriteTCPRequest(clientConn, method, [][]byte{rawKey}, dest, clientSalt, testPayload) + common.Must(err) + + reader, err := ReadTCPResponse(clientConn, method, rawKey, clientSalt) + common.Must(err) + + // The first ReadMultiBuffer drains initialPayload from reader cache + mbInit, err := reader.ReadMultiBuffer() + common.Must(err) + if !bytes.Equal(mbInit[0].Bytes(), testPayload) { + t.Errorf("drained initial payload mismatch: got %s, want %s", mbInit[0].Bytes(), testPayload) + } + buf.ReleaseMulti(mbInit) + + // Send additional stream data + streamData := []byte("stream chunk test") + _ = writer.WriteMultiBuffer(buf.MultiBuffer{buf.FromBytes(streamData)}) + + mb, err := reader.ReadMultiBuffer() + common.Must(err) + if !bytes.Equal(mb[0].Bytes(), streamData) { + t.Errorf("echoed stream data mismatch: got %s, want %s", mb[0].Bytes(), streamData) + } + buf.ReleaseMulti(mb) + }() + + wg.Wait() + + if receivedDest.NetAddr() != dest.NetAddr() { + t.Errorf("destination mismatch: got %s, want %s", receivedDest.NetAddr(), dest.NetAddr()) + } + if diff := cmp.Diff(receivedPayload, testPayload); diff != "" { + t.Errorf("payload mismatch: %s", diff) + } + }) + } +} + +func TestUDPCodec(t *testing.T) { + methods := []string{ + MethodAES128GCM, + MethodAES256GCM, + MethodChaCha20Poly1305, + } + + dest := net.UDPDestination(net.LocalHostIP, net.Port(53)) + payload := []byte("DNS query payload") + + for _, methodName := range methods { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + + psk := make([]byte, method.KeySaltLength) + _, _ = rand.Read(psk) + + clientCodec, err := NewUDPPacketCodec(method, [][]byte{psk}) + common.Must(err) + serverCodec, err := NewUDPServerCodec(method, psk, time.Minute) + common.Must(err) + + session, err := clientCodec.NewClientSession() + common.Must(err) + pktBuf, err := session.EncodePacket(dest, payload) + common.Must(err) + defer pktBuf.Release() + + rawCopy := make([]byte, pktBuf.Len()) + copy(rawCopy, pktBuf.Bytes()) + + decoded, err := serverCodec.DecodePacket(pktBuf.Bytes()) + common.Must(err) + + if decoded.HeaderType != HeaderTypeClient { + t.Errorf("expected header type %d, got %d", HeaderTypeClient, decoded.HeaderType) + } + if decoded.Destination.Port != dest.Port { + t.Errorf("port mismatch: got %d, want %d", decoded.Destination.Port, dest.Port) + } + if !bytes.Equal(decoded.Payload, payload) { + t.Errorf("payload mismatch: got %s, want %s", decoded.Payload, payload) + } + + // Replay same packet wire bytes should fail with ErrPacketIdNotUnique + _, err = serverCodec.DecodePacket(rawCopy) + if err != ErrPacketIdNotUnique { + t.Fatalf("expected ErrPacketIdNotUnique on replay, got: %v", err) + } + }) + } +} + +func TestMultiUserManager(t *testing.T) { + masterKey := generateRandomKey(16) + userKey1 := generateRandomKey(16) + userKey2 := generateRandomKey(16) + + config := &MultiUserServerConfig{ + Method: MethodAES128GCM, + Key: masterKey, + Users: []*protocol.User{ + { + Email: "user1@example.com", + Account: serial.ToTypedMessage(&Account{Key: userKey1}), + }, + }, + } + + inbound, err := NewMultiServer(newTestContext(), config) + common.Must(err) + + if inbound.GetUsersCount(context.Background()) != 1 { + t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background())) + } + + u1 := inbound.GetUser(context.Background(), "user1@example.com") + if u1 == nil || u1.Email != "user1@example.com" { + t.Fatal("user1 not found") + } + + // Add User 2 + rawKey2, _ := base64.StdEncoding.DecodeString(userKey2) + u2 := &protocol.MemoryUser{ + Email: "user2@example.com", + Account: &MemoryAccount{ + Key: rawKey2, + }, + } + err = inbound.AddUser(context.Background(), u2) + common.Must(err) + + if inbound.GetUsersCount(context.Background()) != 2 { + t.Fatalf("expected 2 users, got %d", inbound.GetUsersCount(context.Background())) + } + + // Remove User 1 + err = inbound.RemoveUser(context.Background(), "user1@example.com") + common.Must(err) + + if inbound.GetUsersCount(context.Background()) != 1 { + t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background())) + } + if inbound.GetUser(context.Background(), "user1@example.com") != nil { + t.Fatal("user1 should have been removed") + } +} + +func TestLargeStreamTransfer(t *testing.T) { + method, err := GetCipherMethod(MethodAES128GCM) + common.Must(err) + sessionKey := make([]byte, 16) + _, _ = rand.Read(sessionKey) + + clientAead, err := method.NewAEAD(sessionKey) + common.Must(err) + serverAead, err := method.NewAEAD(sessionKey) + common.Must(err) + + r, w := io.Pipe() + defer r.Close() + defer w.Close() + + writer := NewStreamWriter(w, clientAead) + reader := NewStreamReader(r, serverAead) + + const totalSize = 100 * 1024 // 100 KB + data := make([]byte, totalSize) + _, _ = rand.Read(data) + + errCh := make(chan error, 1) + go func() { + // Write using Write (which splits by MaxPacketSize = 65535) + _, werr := writer.Write(data) + if werr != nil { + errCh <- werr + return + } + _ = w.Close() + errCh <- nil + }() + + var received []byte + for { + mb, rerr := reader.ReadMultiBuffer() + if !mb.IsEmpty() { + for _, b := range mb { + received = append(received, b.Bytes()...) + } + buf.ReleaseMulti(mb) + } + if rerr != nil { + if rerr == io.EOF { + break + } + t.Fatalf("ReadMultiBuffer error: %v", rerr) + } + } + + if werr := <-errCh; werr != nil { + t.Fatalf("writer error: %v", werr) + } + + if len(received) != totalSize { + t.Fatalf("received size mismatch: got %d, want %d", len(received), totalSize) + } + if !bytes.Equal(received, data) { + t.Fatal("received data does not match sent data") + } +} + +func TestClientUDPSessionMultiDestination(t *testing.T) { + for _, methodName := range []string{MethodAES128GCM, MethodAES256GCM, MethodChaCha20Poly1305} { + t.Run(methodName, func(t *testing.T) { + method, err := GetCipherMethod(methodName) + common.Must(err) + rawKey := make([]byte, method.KeySaltLength) + _, _ = rand.Read(rawKey) + + clientCodec, err := NewUDPPacketCodec(method, [][]byte{rawKey}) + common.Must(err) + serverCodec, err := NewUDPServerCodec(method, rawKey, time.Minute) + common.Must(err) + + session, err := clientCodec.NewClientSession() + common.Must(err) + + dest1 := net.UDPDestination(net.LocalHostIP, net.Port(53)) + dest2 := net.UDPDestination(net.IPAddress([]byte{127, 0, 0, 2}), net.Port(53)) + + payload1 := []byte("query-google-dns") + payload2 := []byte("query-cloudflare-dns") + + // Client sends to dest1 and dest2 using SAME session + pkt1, err := session.EncodePacket(dest1, payload1) + common.Must(err) + defer pkt1.Release() + pkt2, err := session.EncodePacket(dest2, payload2) + common.Must(err) + defer pkt2.Release() + + // Server decodes both + dec1, err := serverCodec.DecodePacket(pkt1.Bytes()) + common.Must(err) + dec2, err := serverCodec.DecodePacket(pkt2.Bytes()) + common.Must(err) + + if dec1.SessionID != session.ClientSessionID() || dec2.SessionID != session.ClientSessionID() { + t.Fatalf("both packets must share client session ID %d, got %d and %d", session.ClientSessionID(), dec1.SessionID, dec2.SessionID) + } + if dec1.Destination.String() != dest1.String() { + t.Fatalf("expected dest1 %s, got %s", dest1, dec1.Destination) + } + if dec2.Destination.String() != dest2.String() { + t.Fatalf("expected dest2 %s, got %s", dest2, dec2.Destination) + } + if !bytes.Equal(dec1.Payload, payload1) || !bytes.Equal(dec2.Payload, payload2) { + t.Fatal("payload mismatch") + } + + // Server replies to dest1 and dest2 + respPayload1 := []byte("reply-google-dns") + respPayload2 := []byte("reply-cloudflare-dns") + + respPkt1, err := serverCodec.EncodeServerPacket(dec1.SessionID, dest1, respPayload1) + common.Must(err) + respPkt2, err := serverCodec.EncodeServerPacket(dec2.SessionID, dest2, respPayload2) + common.Must(err) + + // Client decodes replies + clientDec1, err := session.DecodePacket(respPkt1) + common.Must(err) + if clientDec1.Destination.String() != dest1.String() { + t.Fatalf("expected client dec1 dest %s, got %s", dest1, clientDec1.Destination) + } + if !bytes.Equal(clientDec1.Payload, respPayload1) { + t.Fatal("reply payload 1 mismatch") + } + + clientDec2, err := session.DecodePacket(respPkt2) + common.Must(err) + if clientDec2.Destination.String() != dest2.String() { + t.Fatalf("expected client dec2 dest %s, got %s", dest2, clientDec2.Destination) + } + if !bytes.Equal(clientDec2.Payload, respPayload2) { + t.Fatal("reply payload 2 mismatch") + } + }) + } +} diff --git a/proxy/shadowsocks_2022/stream.go b/proxy/shadowsocks_2022/stream.go new file mode 100644 index 000000000..acde16155 --- /dev/null +++ b/proxy/shadowsocks_2022/stream.go @@ -0,0 +1,649 @@ +package shadowsocks_2022 + +import ( + "context" + "crypto/cipher" + "crypto/rand" + "encoding/binary" + "io" + "math" + mrand "math/rand/v2" + "sync" + "time" + + "github.com/xtls/xray-core/common/antireplay" + "github.com/xtls/xray-core/common/buf" + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/protocol" + "github.com/xtls/xray-core/common/signal" + "github.com/xtls/xray-core/common/task" + "github.com/xtls/xray-core/features/policy" + "github.com/xtls/xray-core/transport" +) + +var addrParser = protocol.NewAddressParser( + protocol.AddressFamilyByte(0x01, net.AddressFamilyIPv4), + protocol.AddressFamilyByte(0x04, net.AddressFamilyIPv6), + protocol.AddressFamilyByte(0x03, net.AddressFamilyDomain), + protocol.WithAddressTypeParser(func(b byte) byte { + return b & 0x0F + }), +) + +func IncreaseNonce(nonce []byte) { + for i := range nonce { + nonce[i]++ + if nonce[i] != 0 { + return + } + } +} + +// WriteAddressPort writes a destination address and port in SOCKS5 format +func WriteAddressPort(w io.Writer, dest net.Destination) error { + return addrParser.WriteAddressPort(w, dest.Address, dest.Port) +} + +// AddrPortLength returns the serialized length of a destination in SOCKS5 format +func AddrPortLength(dest net.Destination) int { + switch dest.Address.Family() { + case net.AddressFamilyIPv4: + return 1 + 4 + 2 + case net.AddressFamilyDomain: + return 1 + 1 + len(dest.Address.Domain()) + 2 + case net.AddressFamilyIPv6: + return 1 + 16 + 2 + default: + return 0 + } +} + +type StreamWriter struct { + writer io.Writer + cipher cipher.AEAD + nonce [StreamNonceSize]byte + lenBuf [2]byte + buf []byte +} + +func NewStreamWriter(w io.Writer, c cipher.AEAD) *StreamWriter { + return &StreamWriter{ + writer: w, + cipher: c, + buf: make([]byte, 0, MaxPacketSize+2+2*AEADTagSize), + } +} + +func (w *StreamWriter) Nonce() []byte { + return w.nonce[:] +} + +func (w *StreamWriter) WriteChunk(payload []byte) error { + payloadLen := len(payload) + if payloadLen == 0 { + return nil + } + if payloadLen > MaxPacketSize { + return errors.New("payload exceeds MaxPacketSize") + } + + binary.BigEndian.PutUint16(w.lenBuf[:], uint16(payloadLen)) + w.buf = w.cipher.Seal(w.buf[:0], w.nonce[:], w.lenBuf[:], nil) + IncreaseNonce(w.nonce[:]) + + w.buf = w.cipher.Seal(w.buf, w.nonce[:], payload, nil) + IncreaseNonce(w.nonce[:]) + + _, err := w.writer.Write(w.buf) + return err +} + +func (w *StreamWriter) Write(p []byte) (int, error) { + n := len(p) + for len(p) > 0 { + chunkSize := len(p) + if chunkSize > MaxPacketSize { + chunkSize = MaxPacketSize + } + if err := w.WriteChunk(p[:chunkSize]); err != nil { + return 0, err + } + p = p[chunkSize:] + } + return n, nil +} + +func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { + defer buf.ReleaseMulti(mb) + for _, b := range mb { + p := b.Bytes() + for len(p) > 0 { + chunkSize := len(p) + if chunkSize > MaxPacketSize { + chunkSize = MaxPacketSize + } + if err := w.WriteChunk(p[:chunkSize]); err != nil { + return err + } + p = p[chunkSize:] + } + } + return nil +} + +type StreamReader struct { + reader io.Reader + cipher cipher.AEAD + nonce [StreamNonceSize]byte + lenBuf [2 + AEADTagSize]byte + buffer []byte + cached int + offset int +} + +func NewStreamReader(r io.Reader, c cipher.AEAD) *StreamReader { + return &StreamReader{ + reader: r, + cipher: c, + buffer: make([]byte, MaxPacketSize+AEADTagSize), + } +} + +func (r *StreamReader) Nonce() []byte { + return r.nonce[:] +} + +func (r *StreamReader) Read(p []byte) (int, error) { + if r.cached > 0 { + n := copy(p, r.buffer[r.offset:r.offset+r.cached]) + r.cached -= n + r.offset += n + return n, nil + } + + // Read 2-byte length + AEAD tag (18 bytes) + if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil { + return 0, err + } + + decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil) + if err != nil { + return 0, errors.New("failed to decrypt chunk length").Base(err) + } + IncreaseNonce(r.nonce[:]) + + payloadLen := int(binary.BigEndian.Uint16(decryptedLen)) + if payloadLen == 0 || payloadLen > MaxPacketSize { + return 0, ErrInvalidRequest + } + + chunkEnd := payloadLen + AEADTagSize + if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil { + return 0, err + } + + decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil) + if err != nil { + return 0, errors.New("failed to decrypt chunk payload").Base(err) + } + IncreaseNonce(r.nonce[:]) + + r.cached = len(decryptedPayload) + r.offset = 0 + + n := copy(p, r.buffer[r.offset:r.offset+r.cached]) + r.cached -= n + r.offset += n + return n, nil +} + +func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) { + if r.cached > 0 { + mb := buf.MergeBytes(nil, r.buffer[r.offset:r.offset+r.cached]) + r.cached = 0 + r.offset = 0 + return mb, nil + } + + if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil { + return nil, err + } + + decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil) + if err != nil { + return nil, errors.New("failed to decrypt chunk length").Base(err) + } + IncreaseNonce(r.nonce[:]) + + payloadLen := int(binary.BigEndian.Uint16(decryptedLen)) + if payloadLen == 0 || payloadLen > MaxPacketSize { + return nil, ErrInvalidRequest + } + + chunkEnd := payloadLen + AEADTagSize + if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil { + return nil, err + } + + decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil) + if err != nil { + return nil, errors.New("failed to decrypt chunk payload").Base(err) + } + IncreaseNonce(r.nonce[:]) + + mb := buf.MergeBytes(nil, decryptedPayload) + return mb, nil +} + +type ClientRequestHeader struct { + Destination net.Destination + EarlyData []byte +} + +func ReadClientRequestHeaderWithFixed(reader *StreamReader, fixedChunk []byte) (*ClientRequestHeader, error) { + plainFixed, err := reader.cipher.Open(fixedChunk[:0], reader.Nonce(), fixedChunk, nil) + if err != nil { + return nil, errors.New("failed to decrypt client request header").Base(err) + } + IncreaseNonce(reader.Nonce()) + + if plainFixed[0] != HeaderTypeClient { + return nil, ErrBadHeaderType + } + + epoch := binary.BigEndian.Uint64(plainFixed[1:9]) + diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch)))) + if diff > 30 { + return nil, ErrBadTimestamp + } + + varHeaderLen := int(binary.BigEndian.Uint16(plainFixed[9:11])) + if varHeaderLen == 0 { + return nil, ErrInvalidRequest + } + + var stackVarChunk [512]byte + var varChunkCipher []byte + needed := varHeaderLen + AEADTagSize + if needed <= len(stackVarChunk) { + varChunkCipher = stackVarChunk[:needed] + } else { + varChunkCipher = make([]byte, needed) + } + if _, err := io.ReadFull(reader.reader, varChunkCipher); err != nil { + return nil, err + } + + plainVar, err := reader.cipher.Open(varChunkCipher[:0], reader.Nonce(), varChunkCipher, nil) + if err != nil { + return nil, errors.New("failed to decrypt variable request header").Base(err) + } + IncreaseNonce(reader.Nonce()) + + dest, addrLen, err := ParseAddressPort(plainVar) + if err != nil { + return nil, err + } + dest.Network = net.Network_TCP + + offset := addrLen + if len(plainVar) < offset+2 { + return nil, ErrPacketTooShort + } + paddingLen := int(binary.BigEndian.Uint16(plainVar[offset : offset+2])) + offset += 2 + + if len(plainVar) < offset+paddingLen { + return nil, ErrNoPadding + } + offset += paddingLen + + var earlyData []byte + var payloadLen int + if len(plainVar) > offset { + earlyData = plainVar[offset:] + payloadLen = len(earlyData) + } + + // SIP022 §3.1.4: Servers MUST reject the request if the variable-length header chunk does not contain payload and the padding length is 0. + if paddingLen == 0 && payloadLen == 0 { + return nil, errors.New("request without payload and padding is not allowed") + } + + return &ClientRequestHeader{ + Destination: dest, + EarlyData: earlyData, + }, nil +} + +// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer. +func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) { + finalPSK := pskList[len(pskList)-1] + sessionKey := DeriveSessionSubKey(finalPSK, clientSalt, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + if err != nil { + return nil, err + } + + writer := NewStreamWriter(w, aead) + + payloadLen := len(payload) + var paddingLen int + if payloadLen < MaxPaddingLength { + paddingLen = mrand.IntN(MaxPaddingLength) + 1 + } + addrPortLen := AddrPortLength(dest) + varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen + + totalHandshakeLen := int32(method.KeySaltLength + len(pskList)*AESBlockSize + RequestHeaderFixedChunkLength + AEADTagSize + varHeaderLen + AEADTagSize) + handshakeBuf := buf.NewWithSize(totalHandshakeLen) + defer handshakeBuf.Release() + + handshakeBuf.Write(clientSalt) + + for i, currPSK := range pskList[:len(pskList)-1] { + identitySubkey := DeriveIdentitySubKey(currPSK, clientSalt, method.KeySaltLength) + block, err := method.NewBlock(identitySubkey) + if err != nil { + return nil, err + } + nextPSK := pskList[i+1] + pskHash := DeriveUserPSKHash(nextPSK) + var encryptedEIH [AESBlockSize]byte + block.Encrypt(encryptedEIH[:], pskHash[:]) + handshakeBuf.Write(encryptedEIH[:]) + } + + var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte + fixedHeaderPlaintext[0] = HeaderTypeClient + binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix())) + binary.BigEndian.PutUint16(fixedHeaderPlaintext[9:11], uint16(varHeaderLen)) + + fixedChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedHeaderPlaintext[:], nil) + IncreaseNonce(writer.nonce[:]) + handshakeBuf.Write(fixedChunk) + + varHeaderBuf := buf.NewWithSize(int32(varHeaderLen)) + defer varHeaderBuf.Release() + + if err := WriteAddressPort(varHeaderBuf, dest); err != nil { + return nil, err + } + + var padLenBytes [2]byte + binary.BigEndian.PutUint16(padLenBytes[:], uint16(paddingLen)) + varHeaderBuf.Write(padLenBytes[:]) + + if paddingLen > 0 { + varHeaderBuf.Write(zeroPadding[:paddingLen]) + } + + if payloadLen > 0 { + varHeaderBuf.Write(payload) + } + + varChunk := writer.cipher.Seal(nil, writer.nonce[:], varHeaderBuf.Bytes(), nil) + IncreaseNonce(writer.nonce[:]) + handshakeBuf.Write(varChunk) + + if _, err := w.Write(handshakeBuf.Bytes()); err != nil { + return nil, err + } + + return writer, nil +} + +// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream. +func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) { + fixedPlainLen := 1 + 8 + method.KeySaltLength + 2 + chunkCipherLen := fixedPlainLen + AEADTagSize + headerLen := method.KeySaltLength + chunkCipherLen + + // Single read call for Salt + Fixed-length response header chunk per SIP022 §3.1.4 + var headerBuf [128]byte + headerSlice := headerBuf[:headerLen] + n, err := r.Read(headerSlice) + if err != nil || n < headerLen { + return nil, errors.New("failed to read complete server response header") + } + + serverSaltSlice := headerSlice[:method.KeySaltLength] + chunkSlice := headerSlice[method.KeySaltLength:headerLen] + + sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + if err != nil { + return nil, err + } + + reader := NewStreamReader(r, aead) + + decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil) + if err != nil { + return nil, errors.New("failed to decrypt server response header").Base(err) + } + IncreaseNonce(reader.nonce[:]) + + if decryptedFixed[0] != HeaderTypeServer { + return nil, ErrBadHeaderType + } + + serverEpoch := binary.BigEndian.Uint64(decryptedFixed[1:9]) + diff := int(math.Abs(float64(time.Now().Unix() - int64(serverEpoch)))) + if diff > 30 { + return nil, ErrBadTimestamp + } + + echoedSalt := decryptedFixed[9 : 9+method.KeySaltLength] + for i := 0; i < method.KeySaltLength; i++ { + if echoedSalt[i] != clientSalt[i] { + return nil, errors.New("bad request salt") + } + } + + initialPayloadLen := int(binary.BigEndian.Uint16(decryptedFixed[9+method.KeySaltLength : 11+method.KeySaltLength])) + if initialPayloadLen > 0 { + initialCipherLen := initialPayloadLen + AEADTagSize + if _, err := io.ReadFull(r, reader.buffer[:initialCipherLen]); err != nil { + return nil, err + } + decryptedInitial, err := reader.cipher.Open(reader.buffer[:0], reader.nonce[:], reader.buffer[:initialCipherLen], nil) + if err != nil { + return nil, errors.New("failed to decrypt initial response payload").Base(err) + } + IncreaseNonce(reader.nonce[:]) + reader.cached = len(decryptedInitial) + reader.offset = 0 + } + + return reader, nil +} + +// ServerStreamWriter lazily sends the response header along with the first payload chunk per SIP022 §3.1.2 & §3.1.4. +type ServerStreamWriter struct { + mu sync.Mutex + w io.Writer + method *CipherMethod + psk []byte + clientSalt []byte + streamWriter *StreamWriter +} + +func NewServerStreamWriter(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte) *ServerStreamWriter { + return &ServerStreamWriter{ + w: w, + method: method, + psk: psk, + clientSalt: clientSalt, + } +} + +func (s *ServerStreamWriter) sendHeaderWithFirstPayload(payload []byte) (*StreamWriter, error) { + var serverSalt [32]byte + serverSaltSlice := serverSalt[:s.method.KeySaltLength] + if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil { + return nil, err + } + + respKey := DeriveSessionSubKey(s.psk, serverSaltSlice, s.method.KeySaltLength) + respAead, err := s.method.NewAEAD(respKey) + if err != nil { + return nil, err + } + sw := NewStreamWriter(s.w, respAead) + + totalHeaderLen := int32(s.method.KeySaltLength + 1 + 8 + s.method.KeySaltLength + 2 + AEADTagSize + len(payload) + AEADTagSize) + outBuf := buf.NewWithSize(totalHeaderLen) + defer outBuf.Release() + + outBuf.Write(serverSaltSlice) + + var fixedRespPlain [1 + 8 + 32 + 2]byte + fixedRespSlice := fixedRespPlain[:1+8+s.method.KeySaltLength+2] + fixedRespSlice[0] = HeaderTypeServer + binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix())) + copy(fixedRespSlice[9:9+s.method.KeySaltLength], s.clientSalt) + binary.BigEndian.PutUint16(fixedRespSlice[9+s.method.KeySaltLength:11+s.method.KeySaltLength], uint16(len(payload))) + + fixedRespChunk := sw.cipher.Seal(nil, sw.nonce[:], fixedRespSlice, nil) + IncreaseNonce(sw.nonce[:]) + outBuf.Write(fixedRespChunk) + + if len(payload) > 0 { + payloadChunk := sw.cipher.Seal(nil, sw.nonce[:], payload, nil) + IncreaseNonce(sw.nonce[:]) + outBuf.Write(payloadChunk) + } + + if _, err := s.w.Write(outBuf.Bytes()); err != nil { + return nil, err + } + return sw, nil +} + +func (s *ServerStreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error { + if mb.IsEmpty() { + return nil + } + + if s.streamWriter == nil { + s.mu.Lock() + if s.streamWriter == nil { + firstBuf := mb[0] + firstBytes := firstBuf.Bytes() + chunkSize := len(firstBytes) + if chunkSize > MaxPacketSize { + chunkSize = MaxPacketSize + } + firstPayload := firstBytes[:chunkSize] + sw, err := s.sendHeaderWithFirstPayload(firstPayload) + if err != nil { + s.mu.Unlock() + buf.ReleaseMulti(mb) + return err + } + s.streamWriter = sw + + firstBuf.Advance(int32(chunkSize)) + if firstBuf.IsEmpty() { + firstBuf.Release() + mb = mb[1:] + } + } + s.mu.Unlock() + if len(mb) == 0 { + return nil + } + } + + return s.streamWriter.WriteMultiBuffer(mb) +} + +func (s *ServerStreamWriter) Write(p []byte) (int, error) { + n := len(p) + if s.streamWriter == nil { + s.mu.Lock() + if s.streamWriter == nil { + chunkSize := len(p) + if chunkSize > MaxPacketSize { + chunkSize = MaxPacketSize + } + firstPayload := p[:chunkSize] + sw, err := s.sendHeaderWithFirstPayload(firstPayload) + if err != nil { + s.mu.Unlock() + return 0, err + } + s.streamWriter = sw + p = p[chunkSize:] + } + s.mu.Unlock() + if len(p) == 0 { + return n, nil + } + } + + _, err := s.streamWriter.Write(p) + return n, err +} + +func (s *ServerStreamWriter) Close() error { + if s.streamWriter == nil { + s.mu.Lock() + defer s.mu.Unlock() + if s.streamWriter == nil { + sw, err := s.sendHeaderWithFirstPayload(nil) + if err != nil { + return err + } + s.streamWriter = sw + } + } + return nil +} + +// InitServerStream decrypts the client request header, verifies the timestamp and replay filter, +// and returns a StreamReader for subsequent stream chunks. +func InitServerStream(conn net.Conn, method *CipherMethod, psk, saltSlice []byte, salt [32]byte, fixedChunk []byte, saltFilter *antireplay.ReplayFilter[[32]byte]) (*StreamReader, *ClientRequestHeader, error) { + sessionKey := DeriveSessionSubKey(psk, saltSlice, method.KeySaltLength) + aead, err := method.NewAEAD(sessionKey) + if err != nil { + return nil, nil, err + } + + reader := NewStreamReader(conn, aead) + + reqHeader, err := ReadClientRequestHeaderWithFixed(reader, fixedChunk) + if err != nil { + return nil, nil, err + } + _ = conn.SetReadDeadline(time.Time{}) + + if !saltFilter.Check(salt) { + return nil, nil, ErrSaltNotUnique + } + return reader, reqHeader, nil +} + +func TransportTCP(ctx context.Context, sessionPolicy policy.Session, reader buf.Reader, writer buf.Writer, link *transport.Link) error { + ctx, cancel := context.WithCancel(ctx) + timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle) + ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer) + + requestDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly) + return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)) + } + + responseDone := func() error { + defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) + if c, ok := writer.(io.Closer); ok { + defer c.Close() + } + return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)) + } + + responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer)) + return task.Run(ctx, requestDone, responseDoneAndCloseWriter) +} diff --git a/proxy/socks/client.go b/proxy/socks/client.go index eac3f35b8..aa14f8f69 100644 --- a/proxy/socks/client.go +++ b/proxy/socks/client.go @@ -105,7 +105,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter } udpRequest, err := ClientHandshake(request, conn, conn) if err != nil { - return errors.New("failed to establish connection to server").AtWarning().Base(err) + return errors.New("failed to establish connection to server").Base(err) } if udpRequest != nil { if udpRequest.Address == net.AnyIP || udpRequest.Address == net.AnyIPv6 { diff --git a/proxy/socks/protocol.go b/proxy/socks/protocol.go index 5d4cfd9d9..d71c0fb1e 100644 --- a/proxy/socks/protocol.go +++ b/proxy/socks/protocol.go @@ -458,10 +458,10 @@ func ClientHandshake(request *protocol.RequestHeader, reader io.Reader, writer i } if b.Byte(0) != socks5Version { - return nil, errors.New("unexpected server version: ", b.Byte(0)).AtWarning() + return nil, errors.New("unexpected server version: ", b.Byte(0)) } if b.Byte(1) != authByte { - return nil, errors.New("auth method not supported.").AtWarning() + return nil, errors.New("auth method not supported.") } if authByte == authPassword { diff --git a/proxy/trojan/client.go b/proxy/trojan/client.go index 4af8c019d..abbf3ee7f 100644 --- a/proxy/trojan/client.go +++ b/proxy/trojan/client.go @@ -69,7 +69,7 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter return nil }) if err != nil { - return errors.New("failed to find an available destination").AtWarning().Base(err) + return errors.New("failed to find an available destination").Base(err) } errors.LogInfo(ctx, "tunneling request to ", destination, " via ", server.Destination.NetAddr()) @@ -116,21 +116,21 @@ func (c *Client) Process(ctx context.Context, link *transport.Link, dialer inter // write some request payload to buffer if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout { - return errors.New("failed to write A request payload").Base(err).AtWarning() + return errors.New("failed to write A request payload").Base(err) } // Flush; bufferWriter.WriteMultiBuffer now is bufferWriter.writer.WriteMultiBuffer if err = bufferWriter.SetBuffered(false); err != nil { - return errors.New("failed to flush payload").Base(err).AtWarning() + return errors.New("failed to flush payload").Base(err) } // Send header if not sent yet if _, err = connWriter.Write([]byte{}); err != nil { - return err.(*errors.Error).AtWarning() + return err } if err = buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer)); err != nil { - return errors.New("failed to transfer request payload").Base(err).AtInfo() + return errors.New("failed to transfer request payload").Base(err) } return nil diff --git a/proxy/trojan/server.go b/proxy/trojan/server.go index d5979e529..8fd19beb2 100644 --- a/proxy/trojan/server.go +++ b/proxy/trojan/server.go @@ -47,11 +47,11 @@ func NewServer(ctx context.Context, config *ServerConfig) (*Server, error) { for _, user := range config.Users { u, err := user.ToMemoryUser() if err != nil { - return nil, errors.New("failed to get trojan user").Base(err).AtError() + return nil, errors.New("failed to get trojan user").Base(err) } if err := validator.Add(u); err != nil { - return nil, errors.New("failed to add user").Base(err).AtError() + return nil, errors.New("failed to add user").Base(err) } } @@ -151,7 +151,7 @@ func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Con sessionPolicy := s.policyManager.ForLevel(0) if err := conn.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { - return errors.New("unable to set read deadline").Base(err).AtWarning() + return errors.New("unable to set read deadline").Base(err) } first := buf.FromBytes(make([]byte, buf.Size)) @@ -219,7 +219,7 @@ func (s *Server) Process(ctx context.Context, network net.Network, conn stat.Con destination := clientReader.Target if err := conn.SetReadDeadline(time.Time{}); err != nil { - return errors.New("unable to set read deadline").Base(err).AtWarning() + return errors.New("unable to set read deadline").Base(err) } inbound := session.InboundFromContext(ctx) @@ -402,7 +402,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S } apfb := napfb[name] if apfb == nil { - return errors.New(`failed to find the default "name" config`).AtWarning() + return errors.New(`failed to find the default "name" config`) } if apfb[alpn] == nil { @@ -410,7 +410,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S } pfb := apfb[alpn] if pfb == nil { - return errors.New(`failed to find the default "alpn" config`).AtWarning() + return errors.New(`failed to find the default "alpn" config`) } path := "" @@ -444,7 +444,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S } fb := pfb[path] if fb == nil { - return errors.New(`failed to find the default "path" config`).AtWarning() + return errors.New(`failed to find the default "path" config`) } ctx, cancel := context.WithCancel(ctx) @@ -460,7 +460,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S } return nil }); err != nil { - return errors.New("failed to dial to " + fb.Dest).Base(err).AtWarning() + return errors.New("failed to dial to " + fb.Dest).Base(err) } defer conn.Close() @@ -520,11 +520,11 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S common.Must2(pro.Write([]byte{byte(p1 >> 8), byte(p1), byte(p2 >> 8), byte(p2)})) } if err := serverWriter.WriteMultiBuffer(buf.MultiBuffer{pro}); err != nil { - return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err).AtWarning() + return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err) } } if err := buf.Copy(reader, serverWriter, buf.UpdateActivity(timer)); err != nil { - return errors.New("failed to fallback request payload").Base(err).AtInfo() + return errors.New("failed to fallback request payload").Base(err) } return nil } @@ -534,7 +534,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S getResponse := func() error { defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) if err := buf.Copy(serverReader, writer, buf.UpdateActivity(timer)); err != nil { - return errors.New("failed to deliver response payload").Base(err).AtInfo() + return errors.New("failed to deliver response payload").Base(err) } return nil } @@ -542,7 +542,7 @@ func (s *Server) fallback(ctx context.Context, err error, sessionPolicy policy.S if err := task.Run(ctx, task.OnSuccess(postRequest, task.Close(serverWriter)), task.OnSuccess(getResponse, task.Close(writer))); err != nil { common.Must(common.Interrupt(serverReader)) common.Must(common.Interrupt(serverWriter)) - return errors.New("fallback ends").Base(err).AtInfo() + return errors.New("fallback ends").Base(err) } return nil diff --git a/proxy/tun/README.md b/proxy/tun/README.md index 18c4c3345..70cc66ff6 100644 --- a/proxy/tun/README.md +++ b/proxy/tun/README.md @@ -15,13 +15,57 @@ Plainly enabling it in the config probably will result nothing, or lock your rou ## DETAILS By default, enabling the feature will only bring the tun interface up. \ -When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS uses the first IPv4 prefix from `gateway` to configure the utun point-to-point address. \ +When configured explicitly, Windows and Linux can apply interface addresses from `gateway`, while macOS and FreeBSD use the first IPv4 prefix from `gateway` for the point-to-point address. \ +Without `gateway`, the systems differ: Xray assigns no address on Linux, Windows gives the interface link-local addresses itself (an IPv6 one at once, an IPv4 one from `169.254.0.0/16` after a few seconds), and macOS and FreeBSD use `169.254.10.1/30`. \ Windows, Linux and macOS can also apply system routes from `autoSystemRoutingTable`. -Linux and macOS do not configure system DNS from the `dns` field; system DNS remains managed by the OS or distribution-specific network services. \ +macOS does not configure system DNS from the `dns` field, and neither does Linux by default; system DNS remains managed by the OS or distribution-specific network services. \ For more advanced routing policies or rules, OS level configuration can still manage the named interface (e.g. xray0) when it appears. This keeps complex system level routing and rules in a single place of responsibility - the OS itself. \ Examples of how to achieve this on a simple Linux system (Ubuntu with systemd-networkd) can be found at the end of this README. +### SYSTEM DNS ON LINUX (`autoSystemDnsToGateway`) + +On Linux, setting `autoSystemDnsToGateway` to `true` lets the inbound point the system resolver at the tun interface, so name lookups resolve through Xray instead of going out over the physical link. It is off by default, and it is Linux-only. + +It uses `resolvectl`, which means it only works when all of these hold. Where Xray can tell that one does not, it does not start: + +- the system runs systemd and `resolvectl` is on `PATH` +- `systemd-resolved` is enabled and actually managing DNS (installed but not running is not enough) +- systemd-resolved is version 240 or newer, where `default-route` exists +- no `dns` upstream resolves through the system resolver, directly or through its own bootstrap (see below) + +The address handed over is the first IPv4 `gateway`, or without one the first IPv6 `gateway`, incremented by one (e.g. `192.168.100.1/30` -> `192.168.100.2`, `fc00::1/64` -> `fc00::2`). Without any `gateway`, the config is rejected. It is not taken from `dns`: handing `1.1.1.1` to `resolvectl dns` would make systemd-resolved query that server directly over the physical link, which is the leak this option exists to close. + +Because that address has to actually answer, the takeover is checked before it happens. A query from the interface address to that address is routed through the configured rules, and host-wide DNS is only changed when the result is a DNS-capable outbound. Otherwise DNS is left alone and Xray does not start. In practice this means you also need a routing rule sending the interface's port 53 to a `dns` outbound, for example: + +```json +"routing": { + "rules": [ + { "type": "field", "inboundTag": ["tun"], "port": 53, "outboundTag": "dns" } + ] +} +``` + +The check is a preflight, not a proof for arbitrary rules. It sends its query from the interface address and from a representative ephemeral source port, so a rule that matches on the source port cannot be predicted ahead of time: if the interface's port 53 reaches the `dns` outbound only from some source ports, the takeover is accepted and queries from the other ports fail. Supported configurations are those where the DNS path does not depend on the source port, that is, where the interface's port 53 reaches a `dns` outbound whatever its source. + +It is also a check for the dependencies it knows about, not a proof that no indirect one exists. A hostname-based upstream that bootstraps through system DNS is the case in point: `https+local://dns.google/dns-query` resolves its own hostname with `DialSystem`, so once the takeover is in place that bootstrap goes `resolved -> TUN -> DNS outbound -> bootstrap -> resolved` and the query times out. The preflight does not see it, because the dependency sits in the upstream's bootstrap rather than in the clients it inspects. Upstream resolution, bootstrap included, therefore has to stay independent of the resolver path being redirected; configuring the address instead of the hostname, or resolving the hostname beforehand, avoids it. + +The upstream requirement in the list above matters as much as the routing rule. With no name servers configured, Core resolves through a client that forwards to the system resolver; pointing the system resolver at the TUN would then close a loop through the DNS outbound, `resolved -> TUN -> DNS outbound -> system resolver -> resolved`, and resolution stops. The takeover is refused in that case, and Xray does not start. + +The same applies to a name server pointed at `localhost`, and to a `dns` section that is present but lists no name servers. One such upstream is enough to refuse the takeover even when independent upstreams are configured alongside it: name servers are selected per domain, so a domain-specific rule can still choose the local one, and the loop then affects whichever domains reach it. The check is deliberately broader than the loop it observed, because the alternative would be to drop a name server the user configured. + +Where it cannot apply, Xray does not start, rather than run with the leak described in XTLS/Xray-core#6454, so leave the option off there: + +| Environment | Behaviour | +|---|---| +| systemd distribution with systemd-resolved enabled | applies | +| Alpine, Void, Devuan, OpenRC-based, OpenWrt | no `resolvectl`, does not start | +| DNS managed by dnsmasq / unbound / BIND / static `resolv.conf` | unreachable by `resolvectl`, does not start | +| Containers without a systemd-resolved daemon | does not start | +| systemd older than 240 | `default-route` unavailable, does not start | + +On `Close()` the setting is reverted. It is **not** reverted if the process is killed with `SIGKILL`, since a process cannot handle that signal; run `resolvectl revert ` to clean up by hand. An application that brings its own DNS endpoint is unaffected either way — this only covers the system resolver. + Due to this inbound not actually being a proxy, the configuration ignore required listen and port options, and never listen on any port. \ Here is simple Xray config snippet to enable the inbound: ``` @@ -155,6 +199,20 @@ To make it start, wintun.dll specific for your Windows/arch must be present next After the start network adapter with the name you chose in the config will be created in the system, and exist while Xray is running. +When `dns` is set, those servers are applied to the adapter. Windows is kept from registering the TUN's addresses in DNS, and its DNS cache is flushed when the TUN starts and stops. + +With `autoSystemWfpBlockLeak`, which needs `autoSystemRoutingTable` (the config is rejected otherwise), Xray also adds Windows Filtering Platform filters that keep two kinds of traffic of every program but Xray itself from leaving outside the TUN, each chosen by a value in the list, e.g. `"autoSystemWfpBlockLeak": ["dns", "misconfigtun"]`: +- `"dns"` (needs `dns`, the config is rejected otherwise): DNS (port 53) only goes through the TUN. Windows keeps sending name queries to the DNS servers of the other interfaces as well, out through those interfaces whatever the routes say, and other programs reach a resolver on the local network (e.g. `192.168.1.1` handed out by DHCP) through its more specific LAN route instead of the TUN. On Windows 11 and Server 2022 and later, where those queries may also go over HTTPS or TLS, Windows' DNS Client service cannot connect outside the TUN at all, except for name resolution on the local network (LLMNR, mDNS). The `dns` servers therefore have to lie within `gateway` or `autoSystemRoutingTable` (a warning is logged otherwise), and DNS servers that should be reached directly belong in Xray's own `dns` settings. +- `"misconfigtun"`: an IP version without routes in `autoSystemRoutingTable`, IPv4 or IPv6, is blocked entirely, in both directions, as it would bypass the TUN. Only loopback and what Windows itself needs on the local link (DHCP, and for IPv6 neighbor and multicast listener discovery) remain allowed. An address of that version in `gateway` is not needed: without one, Windows gives the TUN link-local addresses itself, an IPv6 one at once and an IPv4 one from `169.254.0.0/16` after some seconds (until then, IPv4 routed to the TUN is unreachable), and what is routed to the TUN goes through it with those. + +With the filters in place, Xray's own connections out also get past Windows Firewall's block rules (other firewalls may still block them), while connections to Xray's inbounds stay subject to them. + +Names that Xray resolves through the system resolver, such as an outbound's server address given as a domain with the default `AsIs` domain strategy, would be looked up by Windows on Xray's behalf, and those queries would then go into the TUN too. While DNS is restricted this way and `autoOutboundsInterface` is in use (the default with `autoSystemRoutingTable`), Xray therefore resolves them itself, with its own queries to the DNS servers of the other interfaces. That bypasses Windows' DNS cache, and its name resolution on the local network (LLMNR, mDNS): a server address given as a domain is looked up again for every connection, and a DNS server that does not answer delays each lookup. Having Xray's own `dns` resolve it, through the outbound's `sockopt.domainStrategy`, avoids that. The `localhost` DNS server queries the same servers whenever `autoOutboundsInterface` is in use. Both skip the TUN's own DNS servers, unless another interface uses them as well: queried from Xray itself, they would lead back into it, or nowhere. + +If the filters cannot be added, Xray does not start. They are removed when Xray exits. Not covered is name resolution on the local network (LLMNR, mDNS, NetBIOS), except over an IP version that is blocked. + +`autoSystemWfpBlockLeak` (Windows only) is empty by default, as the filters break some setups: with `"dns"`, a local DNS resolver other programs use (e.g. on `127.0.0.1:53`), the DNS of another VPN on its own interface, virtual machines whose NAT resolves names on the host, or signing in to a captive portal; with `"misconfigtun"`, IPv4 or IPv6 on the local network while no route of that version leads to the TUN. Without the filters, DNS may leak as described above. To keep an IP version out of the TUN on purpose while still blocking DNS leaks, use only `["dns"]`. + You can give the adapter ip address manually, you can live Windows to give it autogenerated ip address (which take few seconds), it doesn't matter, the traffic going _through_ the interface will be forwarded into the app for proxying. \ Minimal configuration that will work for local machine is routing passing the traffic on-link through the interface. You will need the interface id for that, unfortunately it is going to change with every Xray start due to implementation ambiguity between Xray and wintun driver. diff --git a/proxy/tun/config.pb.go b/proxy/tun/config.pb.go index 33bc2ba6f..6a862a187 100644 --- a/proxy/tun/config.pb.go +++ b/proxy/tun/config.pb.go @@ -32,6 +32,8 @@ type Config struct { AutoSystemRoutingTable []string `protobuf:"bytes,6,rep,name=auto_system_routing_table,json=autoSystemRoutingTable,proto3" json:"auto_system_routing_table,omitempty"` AutoOutboundsInterface string `protobuf:"bytes,7,opt,name=auto_outbounds_interface,json=autoOutboundsInterface,proto3" json:"auto_outbounds_interface,omitempty"` Desc string `protobuf:"bytes,8,opt,name=desc,proto3" json:"desc,omitempty"` + AutoSystemDnsToGateway bool `protobuf:"varint,9,opt,name=auto_system_dns_to_gateway,json=autoSystemDnsToGateway,proto3" json:"auto_system_dns_to_gateway,omitempty"` + AutoSystemWfpBlockLeak []string `protobuf:"bytes,10,rep,name=auto_system_wfp_block_leak,json=autoSystemWfpBlockLeak,proto3" json:"auto_system_wfp_block_leak,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } @@ -122,11 +124,25 @@ func (x *Config) GetDesc() string { return "" } +func (x *Config) GetAutoSystemDnsToGateway() bool { + if x != nil { + return x.AutoSystemDnsToGateway + } + return false +} + +func (x *Config) GetAutoSystemWfpBlockLeak() []string { + if x != nil { + return x.AutoSystemWfpBlockLeak + } + return nil +} + var File_proxy_tun_config_proto protoreflect.FileDescriptor const file_proxy_tun_config_proto_rawDesc = "" + "\n" + - "\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\x82\x02\n" + + "\x16proxy/tun/config.proto\x12\x0exray.proxy.tun\"\xfa\x02\n" + "\x06Config\x12\x12\n" + "\x04name\x18\x01 \x01(\tR\x04name\x12\x10\n" + "\x03MTU\x18\x02 \x01(\rR\x03MTU\x12\x18\n" + @@ -136,7 +152,10 @@ const file_proxy_tun_config_proto_rawDesc = "" + "user_level\x18\x05 \x01(\rR\tuserLevel\x129\n" + "\x19auto_system_routing_table\x18\x06 \x03(\tR\x16autoSystemRoutingTable\x128\n" + "\x18auto_outbounds_interface\x18\a \x01(\tR\x16autoOutboundsInterface\x12\x12\n" + - "\x04desc\x18\b \x01(\tR\x04descBL\n" + + "\x04desc\x18\b \x01(\tR\x04desc\x12:\n" + + "\x1aauto_system_dns_to_gateway\x18\t \x01(\bR\x16autoSystemDnsToGateway\x12:\n" + + "\x1aauto_system_wfp_block_leak\x18\n" + + " \x03(\tR\x16autoSystemWfpBlockLeakBL\n" + "\x12com.xray.proxy.tunP\x01Z#github.com/xtls/xray-core/proxy/tun\xaa\x02\x0eXray.Proxy.Tunb\x06proto3" var ( diff --git a/proxy/tun/config.proto b/proxy/tun/config.proto index 376ac5af4..a6dfbfcab 100644 --- a/proxy/tun/config.proto +++ b/proxy/tun/config.proto @@ -15,4 +15,6 @@ message Config { repeated string auto_system_routing_table = 6; string auto_outbounds_interface = 7; string desc = 8; + bool auto_system_dns_to_gateway = 9; + repeated string auto_system_wfp_block_leak = 10; } diff --git a/proxy/tun/handler.go b/proxy/tun/handler.go index 53c74a5e3..c13625b0b 100644 --- a/proxy/tun/handler.go +++ b/proxy/tun/handler.go @@ -165,6 +165,18 @@ func (t *Handler) Start() error { return err } + // Platform-specific system DNS takeover, where the platform implements it. + // Rather no TUN than one that the system DNS bypasses. + if c, ok := tunInterface.(interface { + ConfigureSystemDNS(context.Context, string) error + }); ok { + if err := c.ConfigureSystemDNS(t.ctx, t.tag); err != nil { + _ = tunStack.Close() + _ = tunInterface.Close() + return errors.New("unable to set the system DNS (remove autoSystemDnsToGateway to run without)").Base(err) + } + } + t.stack = tunStack t.tun = tunInterface diff --git a/proxy/tun/tun_linux.go b/proxy/tun/tun_linux.go index b2c5d35a1..78e2dd400 100644 --- a/proxy/tun/tun_linux.go +++ b/proxy/tun/tun_linux.go @@ -6,12 +6,24 @@ import ( "context" "net" "net/netip" + "os/exec" "strconv" "sync" "github.com/vishvananda/netlink" + appdns "github.com/xtls/xray-core/app/dns" "github.com/xtls/xray-core/common/errors" + xnet "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/platform" + "github.com/xtls/xray-core/common/serial" + "github.com/xtls/xray-core/common/session" + "github.com/xtls/xray-core/core" + feature_dns "github.com/xtls/xray-core/features/dns" + "github.com/xtls/xray-core/features/dns/localdns" + "github.com/xtls/xray-core/features/outbound" + "github.com/xtls/xray-core/features/routing" + routingsession "github.com/xtls/xray-core/features/routing/session" + "github.com/xtls/xray-core/proxy/dns" "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/tcpip/link/fdbased" "gvisor.dev/gvisor/pkg/tcpip/stack" @@ -30,6 +42,244 @@ type LinuxTun struct { systemRoutes []netlink.Route routeMonitorStop chan struct{} routeMonitorOnce sync.Once + + systemDNSSet bool + systemDNSDirty bool +} + +// resolvectlRunner runs a resolvectl command. Overridable for tests. +var resolvectlRunner = func(name string, args ...string) ([]byte, error) { + return exec.Command(name, args...).CombinedOutput() +} + +// systemDNSAddrs derives the addresses used for the system DNS takeover from the +// first IPv4 gateway, or without one, the first IPv6 gateway: the gateway +// address itself is what a query from this interface appears to come from, and +// the next address is what the resolver is pointed at. The latter belongs to +// the TUN and is answered inside Xray; handing the configured public resolvers +// to resolvectl instead would leave the system querying them directly over the +// physical link, defeating the point of the TUN. +func systemDNSAddrs(gateway []string) (source, dns netip.Addr, ok bool) { + var first6 netip.Addr + for _, address := range gateway { + prefix, err := netip.ParsePrefix(address) + if err != nil { + continue + } + addr := prefix.Addr() + if addr.Is4() { + return addr, addr.Next(), true + } + if !first6.IsValid() { + first6 = addr + } + } + if first6.IsValid() { + return first6, first6.Next(), true + } + return netip.Addr{}, netip.Addr{}, false +} + +func buildResolvectlArgs(action, iface string, extra ...string) []string { + args := make([]string, 0, 2+len(extra)) + args = append(args, action, iface) + args = append(args, extra...) + return args +} + +func runResolvectl(action, iface string, extra ...string) error { + args := buildResolvectlArgs(action, iface, extra...) + if _, err := resolvectlRunner("resolvectl", args...); err != nil { + return errors.New("resolvectl ", action, " failed").Base(err) + } + return nil +} + +// ifaceName returns the TUN interface name, or empty when the link is not +// available. Callers must treat empty as "nothing to configure". +func (t *LinuxTun) ifaceName() string { + if t.tunLink == nil { + return "" + } + attrs := t.tunLink.Attrs() + if attrs == nil { + return "" + } + return attrs.Name +} + +// probeSourcePort is a representative client port for the routing probe. A real +// query arrives from an ephemeral port that cannot be known in advance, so this +// only matters for a rule that matches on a source port. +const probeSourcePort = 49152 + +// verifyDNSRouting reports whether a DNS query to address would actually be +// handled. Redirecting the system resolver at an address nothing answers would +// break name resolution outright, so the takeover only proceeds when routing +// hands such a query to a DNS-capable outbound. +// +// Overridable for tests. +var verifyDNSRouting = func(ctx context.Context, inboundTag, source, address string) error { + ip, err := netip.ParseAddr(address) + if err != nil { + return errors.New("invalid DNS address ", address).Base(err) + } + src, err := netip.ParseAddr(source) + if err != nil || src.Is4() != ip.Is4() { + return errors.New("invalid source address ", source).Base(err) + } + + instance := core.MustFromContext(ctx) + + // Any resolution path that could still reach the system resolver has to be + // refused, because pointing the system resolver at the TUN would close a + // loop through the DNS outbound. With no `dns` section Core installs such a + // client; with a `dns` section that has no name servers app/dns falls back + // to one; and a name server pointed at "localhost" is one even when + // independent upstreams are configured alongside it, because name servers + // are selected per domain. + switch dnsFeature := instance.GetFeature(feature_dns.ClientType()).(type) { + case *localdns.Client: + return errors.New("DNS feature is the system resolver, takeover would loop") + case *appdns.DNS: + if dnsFeature.MayUseSystemResolver() { + return errors.New("DNS configuration may resolve through the system resolver, takeover would loop") + } + } + + router, ok := instance.GetFeature(routing.RouterType()).(routing.Router) + if !ok { + return errors.New("router feature unavailable") + } + + // A real query from this interface carries a source address, and rules may + // match on it, so the probe has to carry one too. + queryCtx := session.ContextWithInbound(ctx, &session.Inbound{ + Name: "tun", + Tag: inboundTag, + Source: xnet.UDPDestination(xnet.IPAddress(src.AsSlice()), probeSourcePort), + }) + queryCtx = session.ContextWithOutbounds(queryCtx, []*session.Outbound{{ + Target: xnet.UDPDestination(xnet.IPAddress(ip.AsSlice()), 53), + }}) + + route, err := router.PickRoute(routingsession.AsRoutingContext(queryCtx)) + if err != nil { + return errors.New("no route for ", address, ":53").Base(err) + } + + manager, ok := instance.GetFeature(outbound.ManagerType()).(outbound.Manager) + if !ok { + return errors.New("outbound manager unavailable") + } + + handler := manager.GetHandler(route.GetOutboundTag()) + if handler == nil { + return errors.New("outbound ", route.GetOutboundTag(), " does not exist") + } + if settings := handler.ProxySettings(); settings == nil || settings.Type != serial.GetMessageType(&dns.Config{}) { + return errors.New("outbound ", route.GetOutboundTag(), " does not handle DNS") + } + return nil +} + +// ConfigureSystemDNS points systemd-resolved at this interface so name lookups +// resolve through Xray instead of leaking to the physical link. +// +// It acts only when the config opts in, and it verifies the data path first: +// unless a query to the advertised address would actually be handled, host-wide +// resolution is left to the OS and an error returned. The caller does not start +// the TUN on an error, as the system DNS would bypass it. +func (t *LinuxTun) ConfigureSystemDNS(ctx context.Context, inboundTag string) error { + if !t.options.AutoSystemDnsToGateway { + return nil + } + if t.systemDNSSet { + return nil + } + + // A previous revert may have failed. Retry before applying anything, so a + // dirty resolver does not silently outlive the attempt to clean it up. + if t.systemDNSDirty { + if err := t.revertSystemDNS(); err != nil { + return errors.New("previous system DNS revert still failing").Base(err) + } + } + + source, address, ok := systemDNSAddrs(t.options.Gateway) + if !ok { + return errors.New("no gateway, cannot derive a system DNS address") + } + + iface := t.ifaceName() + if iface == "" { + return errors.New("interface not available") + } + + if err := verifyDNSRouting(ctx, inboundTag, source.String(), address.String()); err != nil { + return errors.New("no DNS path at ", address.String(), ":53").Base(err) + } + + // Applied as a sequence with rollback: a half-configured resolver would be + // worse than none at all. + if err := runResolvectl("dns", iface, address.String()); err != nil { + return errors.New("resolvectl dns failed").Base(err) + } + if err := runResolvectl("domain", iface, "~."); err != nil { + return t.rollbackSystemDNS(iface, errors.New("resolvectl domain failed").Base(err)) + } + if err := runResolvectl("default-route", iface, "true"); err != nil { + return t.rollbackSystemDNS(iface, errors.New("resolvectl default-route failed").Base(err)) + } + + t.systemDNSSet = true + errors.LogInfo(ctx, "[tun] system DNS set to ", address.String(), " on ", iface) + return nil +} + +// rollbackSystemDNS undoes a partially applied takeover. A failed revert is +// recorded so the next attempt retries it, and is reported rather than +// swallowed. +func (t *LinuxTun) rollbackSystemDNS(iface string, cause error) error { + if err := runResolvectl("revert", iface); err != nil { + t.systemDNSDirty = true + // Combine, because Base overwrites: reporting only the cause would hide + // the revert failure, and reporting only the revert failure would hide + // why the revert was attempted. + return errors.New("revert failed, per-link DNS settings may remain").Base(errors.Combine(err, cause)) + } + return cause +} + +// revertSystemDNS issues the revert and keeps the dirty flag in step with the +// outcome. +func (t *LinuxTun) revertSystemDNS() error { + err := runResolvectl("revert", t.ifaceName()) + t.systemDNSDirty = err != nil + if err != nil { + return err + } + t.systemDNSSet = false + return nil +} + +// unsetSystemDNS hands DNS back to the OS. Only meaningful when +// ConfigureSystemDNS applied something, or a previous revert failed. +func (t *LinuxTun) unsetSystemDNS() { + if !t.systemDNSSet && !t.systemDNSDirty { + return + } + + if t.ifaceName() == "" { + // The link is gone, and its per-link settings went with it. + t.systemDNSSet = false + t.systemDNSDirty = false + return + } + + if err := t.revertSystemDNS(); err != nil { + errors.LogInfoInner(context.Background(), err, "[tun] failed to revert system DNS; per-link settings may remain until revert succeeds") + } } // LinuxTun implements Tun @@ -200,6 +450,7 @@ func (t *LinuxTun) Close() error { } }) + t.unsetSystemDNS() _ = t.unsetSystemRoutes() _ = t.unsetInterfaceAddresses() diff --git a/proxy/tun/tun_linux_dns_route_test.go b/proxy/tun/tun_linux_dns_route_test.go new file mode 100644 index 000000000..120ef7e10 --- /dev/null +++ b/proxy/tun/tun_linux_dns_route_test.go @@ -0,0 +1,205 @@ +//go:build linux && !android + +package tun + +import ( + "context" + "strings" + "testing" + + "github.com/xtls/xray-core/app/dispatcher" + appdns "github.com/xtls/xray-core/app/dns" + "github.com/xtls/xray-core/app/proxyman" + _ "github.com/xtls/xray-core/app/proxyman/inbound" + _ "github.com/xtls/xray-core/app/proxyman/outbound" + "github.com/xtls/xray-core/app/router" + "github.com/xtls/xray-core/common/geodata" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/serial" + "github.com/xtls/xray-core/core" + "github.com/xtls/xray-core/proxy/blackhole" + proxydns "github.com/xtls/xray-core/proxy/dns" + "github.com/xtls/xray-core/proxy/freedom" +) + +const ( + routeTestInboundTag = "tun" + routeTestSource = "192.168.100.1" + routeTestDNSAddress = "192.168.100.2" +) + +// port53Rule sends DNS queries arriving from the interface to the dns outbound. +func port53Rule() *router.RoutingRule { + return &router.RoutingRule{ + InboundTag: []string{routeTestInboundTag}, + PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(53)}}, + TargetTag: &router.RoutingRule_Tag{Tag: "dns"}, + } +} + +// sourceBlockRule diverts traffic from one address, which is the shape of a rule +// that only matches because the real request carries a source. +func sourceBlockRule(ip []byte) *router.RoutingRule { + return &router.RoutingRule{ + SourceIp: []*geodata.IPRule{{ + Value: &geodata.IPRule_Custom{ + Custom: &geodata.CIDRRule{ + Cidr: &geodata.CIDR{Ip: ip, Prefix: 32}, + }, + }, + }}, + TargetTag: &router.RoutingRule_Tag{Tag: "block"}, + } +} + +// newRouteTestContext builds a real but unstarted instance: no TUN device, no +// running resolver. The instance is placed in the context through the key core +// exports for tests. +func newRouteTestContext(t *testing.T, withDNSApp bool, nameServers []*appdns.NameServer, rules []*router.RoutingRule) context.Context { + t.Helper() + + apps := []*serial.TypedMessage{ + serial.ToTypedMessage(&dispatcher.Config{}), + serial.ToTypedMessage(&proxyman.InboundConfig{}), + serial.ToTypedMessage(&proxyman.OutboundConfig{}), + serial.ToTypedMessage(&router.Config{Rule: rules}), + } + if withDNSApp { + apps = append(apps, serial.ToTypedMessage(&appdns.Config{NameServer: nameServers})) + } + + instance, err := core.New(&core.Config{ + App: apps, + Outbound: []*core.OutboundHandlerConfig{ + {Tag: "direct", ProxySettings: serial.ToTypedMessage(&freedom.Config{})}, + {Tag: "dns", ProxySettings: serial.ToTypedMessage(&proxydns.Config{})}, + {Tag: "block", ProxySettings: serial.ToTypedMessage(&blackhole.Config{})}, + }, + }) + if err != nil { + t.Fatalf("core.New: %v", err) + } + t.Cleanup(func() { _ = instance.Close() }) + + return context.WithValue(context.Background(), core.XrayKey(1), instance) +} + +func udpNameServer(ip []byte) []*appdns.NameServer { + return []*appdns.NameServer{{ + Address: &net.Endpoint{ + Network: net.Network_UDP, + Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: ip}}, + Port: 53, + }, + }} +} + +// localNameServer is a name server pointed at "localhost", which app/dns +// resolves through the system resolver. +func localNameServer() *appdns.NameServer { + return &appdns.NameServer{ + Address: &net.Endpoint{ + Network: net.Network_UDP, + Address: &net.IPOrDomain{Address: &net.IPOrDomain_Domain{Domain: "localhost"}}, + Port: 53, + }, + } +} + +// These drive the real feature lookup and the real router. verifyDNSRouting is +// the same function ConfigureSystemDNS calls, so a false positive here is a +// false positive in the takeover decision itself, which is what assertions on +// the resolvectl arguments could never catch. +func TestVerifyDNSRoutingDecisions(t *testing.T) { + tests := []struct { + name string + withDNSApp bool + nameServers []*appdns.NameServer + rules []*router.RoutingRule + wantErr string + }{ + { + name: "independent upstream reaches the dns outbound", + withDNSApp: true, + nameServers: udpNameServer([]byte{9, 9, 9, 9}), + rules: []*router.RoutingRule{port53Rule()}, + wantErr: "", + }, + { + name: "no dns section falls back to the system resolver", + rules: []*router.RoutingRule{port53Rule()}, + wantErr: "system resolver", + }, + { + name: "dns section without name servers falls back too", + withDNSApp: true, + rules: []*router.RoutingRule{port53Rule()}, + wantErr: "system resolver", + }, + { + // An independent upstream is not enough on its own: name servers are + // selected per domain, so a local one can still be the one chosen. + // The refusal is deliberately domain-agnostic for that reason. + name: "a local name server alongside an independent one", + withDNSApp: true, + nameServers: append(udpNameServer([]byte{9, 9, 9, 9}), localNameServer()), + rules: []*router.RoutingRule{port53Rule()}, + wantErr: "system resolver", + }, + { + name: "a rule on the interface address diverts the real query", + withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}), + rules: []*router.RoutingRule{ + sourceBlockRule([]byte{192, 168, 100, 1}), + port53Rule(), + }, + wantErr: "does not handle DNS", + }, + { + name: "a rule on another address does not match it", + withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}), + rules: []*router.RoutingRule{ + sourceBlockRule([]byte{10, 0, 0, 1}), + port53Rule(), + }, + wantErr: "", + }, + { + name: "no rule matches the query", + withDNSApp: true, nameServers: udpNameServer([]byte{9, 9, 9, 9}), + wantErr: "no route", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ctx := newRouteTestContext(t, tt.withDNSApp, tt.nameServers, tt.rules) + err := verifyDNSRouting(ctx, routeTestInboundTag, routeTestSource, routeTestDNSAddress) + + if tt.wantErr == "" { + if err != nil { + t.Fatalf("expected the takeover to be accepted, got: %v", err) + } + return + } + if err == nil { + t.Fatalf("expected the takeover to be refused with %q, got nil", tt.wantErr) + } + if !strings.Contains(err.Error(), tt.wantErr) { + t.Errorf("error = %q, want it to contain %q", err.Error(), tt.wantErr) + } + }) + } +} + +// Without an IPv4 gateway, the takeover uses the first IPv6 one, and the probe +// carries IPv6 addresses. +func TestVerifyDNSRoutingIPv6(t *testing.T) { + ctx := newRouteTestContext(t, true, udpNameServer([]byte{9, 9, 9, 9}), []*router.RoutingRule{port53Rule()}) + if err := verifyDNSRouting(ctx, routeTestInboundTag, "fc00::1", "fc00::2"); err != nil { + t.Fatalf("expected the takeover to be accepted, got: %v", err) + } + if err := verifyDNSRouting(ctx, routeTestInboundTag, routeTestSource, "fc00::2"); err == nil { + t.Fatal("expected mixed IPv4 and IPv6 addresses to be refused") + } +} diff --git a/proxy/tun/tun_linux_dns_test.go b/proxy/tun/tun_linux_dns_test.go new file mode 100644 index 000000000..d1c2fdfa7 --- /dev/null +++ b/proxy/tun/tun_linux_dns_test.go @@ -0,0 +1,445 @@ +//go:build linux && !android + +package tun + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/vishvananda/netlink" +) + +// testLink returns a minimal netlink.Link whose Attrs().Name is name, so the +// DNS helpers can be exercised without a real TUN device. +func testLink(name string) netlink.Link { + return &netlink.Dummy{LinkAttrs: netlink.LinkAttrs{Name: name}} +} + +type probeCall struct { + inboundTag string + source string + address string +} + +// stubDNSRouting replaces the routing probe for the duration of a test and +// records how it was called, so tests can assert the probe is representative. +func stubDNSRouting(t *testing.T, err error) *[]probeCall { + t.Helper() + original := verifyDNSRouting + calls := []probeCall{} + verifyDNSRouting = func(_ context.Context, inboundTag, source, address string) error { + calls = append(calls, probeCall{inboundTag, source, address}) + return err + } + t.Cleanup(func() { verifyDNSRouting = original }) + return &calls +} + +// recorder installs a resolvectl stub for the duration of a test and returns the +// captured invocations. An empty failOn succeeds every call; otherwise the named +// subcommand fails. +func recorder(t *testing.T, failOn string) *[][]string { + t.Helper() + original := resolvectlRunner + calls := [][]string{} + resolvectlRunner = func(name string, args ...string) ([]byte, error) { + calls = append(calls, append([]string{name}, args...)) + if failOn != "" && len(args) > 0 && args[0] == failOn { + return nil, errors.New("boom") + } + return nil, nil + } + t.Cleanup(func() { resolvectlRunner = original }) + return &calls +} + +func optedInTun() *LinuxTun { + return &LinuxTun{ + options: &Config{ + Name: "xray_tun", + Gateway: []string{"192.168.100.1/30"}, + AutoSystemDnsToGateway: true, + }, + tunLink: testLink("xray_tun"), + } +} + +func joined(calls [][]string) string { + parts := make([]string, 0, len(calls)) + for _, call := range calls { + parts = append(parts, strings.Join(call, " ")) + } + return strings.Join(parts, " | ") +} + +func TestConfigureSystemDNSDisabledByDefault(t *testing.T) { + probes := stubDNSRouting(t, nil) + calls := recorder(t, "") + + t1 := optedInTun() + t1.options.AutoSystemDnsToGateway = false + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(*probes) != 0 { + t.Errorf("routing probe must not run when disabled, got %d calls", len(*probes)) + } + if len(*calls) != 0 { + t.Errorf("resolvectl must not run when disabled, got %v", *calls) + } + if t1.systemDNSSet { + t.Error("systemDNSSet should stay false when disabled") + } +} + +func TestConfigureSystemDNSNoGateway(t *testing.T) { + probes := stubDNSRouting(t, nil) + calls := recorder(t, "") + + t1 := optedInTun() + t1.options.Gateway = nil + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil { + t.Fatal("expected an error when no gateway is configured") + } + if len(*probes) != 0 { + t.Errorf("routing probe must not run without a gateway, got %d calls", len(*probes)) + } + if len(*calls) != 0 { + t.Errorf("resolvectl must not run without a gateway, got %v", *calls) + } +} + +// This is the case the reviewer flagged: without a routed DNS path, pointing the +// system resolver at the derived address would break resolution outright. +func TestConfigureSystemDNSLeavesOSDNSWhenNoRoute(t *testing.T) { + probes := stubDNSRouting(t, errors.New("no route")) + calls := recorder(t, "") + + t1 := optedInTun() + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil { + t.Fatal("expected an error when the DNS path is unverified") + } + if len(*probes) != 1 { + t.Errorf("routing probe should run once, got %d", len(*probes)) + } + if len(*calls) != 0 { + t.Errorf("system DNS must be left untouched, got %v", *calls) + } + if t1.systemDNSSet { + t.Error("systemDNSSet should stay false when the path is unverified") + } +} + +// A real query from the interface carries a source address, and rules may match +// on it, so the probe must not be source-less. +func TestConfigureSystemDNSProbeCarriesSource(t *testing.T) { + probes := stubDNSRouting(t, nil) + recorder(t, "") + + t1 := optedInTun() + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(*probes) != 1 { + t.Fatalf("expected one probe call, got %d", len(*probes)) + } + got := (*probes)[0] + if got.source != "192.168.100.1" { + t.Errorf("probe source = %q, want the interface address %q", got.source, "192.168.100.1") + } + if got.address != "192.168.100.2" { + t.Errorf("probe address = %q, want %q", got.address, "192.168.100.2") + } + if got.inboundTag != "tun" { + t.Errorf("probe inbound tag = %q, want %q", got.inboundTag, "tun") + } +} + +func TestConfigureSystemDNSAppliesResolvectl(t *testing.T) { + stubDNSRouting(t, nil) + calls := recorder(t, "") + + t1 := optedInTun() + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !t1.systemDNSSet { + t.Fatal("systemDNSSet should be true after a successful takeover") + } + + want := "resolvectl dns xray_tun 192.168.100.2 | " + + "resolvectl domain xray_tun ~. | " + + "resolvectl default-route xray_tun true" + if got := joined(*calls); got != want { + t.Errorf("resolvectl calls = %q, want %q", got, want) + } +} + +func TestConfigureSystemDNSIdempotent(t *testing.T) { + stubDNSRouting(t, nil) + calls := recorder(t, "") + + t1 := optedInTun() + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil { + t.Fatalf("unexpected error: %v", err) + } + before := len(*calls) + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(*calls) != before { + t.Errorf("second call must be a no-op, calls went %d -> %d", before, len(*calls)) + } +} + +// A half-applied resolver is worse than none, so a failure mid-sequence reverts. +func TestConfigureSystemDNSRollsBackOnPartialFailure(t *testing.T) { + stubDNSRouting(t, nil) + calls := recorder(t, "domain") + + t1 := optedInTun() + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil { + t.Fatal("expected an error when a resolvectl step fails") + } + if t1.systemDNSSet { + t.Error("systemDNSSet should stay false after a failed takeover") + } + if t1.systemDNSDirty { + t.Error("a successful revert should not leave the resolver dirty") + } + if !strings.Contains(joined(*calls), "resolvectl revert xray_tun") { + t.Errorf("expected a revert after partial failure, got %q", joined(*calls)) + } +} + +// If the revert itself fails the settings may still be installed, so the state +// has to be remembered rather than silently dropped. +func TestConfigureSystemDNSRollbackFailureKeepsDirty(t *testing.T) { + stubDNSRouting(t, nil) + calls := recorder(t, "revert") + + t1 := optedInTun() + t1.options.Gateway = []string{"192.168.100.1/30"} + // Make only the rollback path fail: "dns" succeeds, "domain" fails, "revert" fails. + *calls = nil + + original := resolvectlRunner + defer func() { resolvectlRunner = original }() + resolvectlRunner = func(name string, args ...string) ([]byte, error) { + *calls = append(*calls, append([]string{name}, args...)) + if len(args) > 0 && (args[0] == "domain" || args[0] == "revert") { + return nil, errors.New("boom") + } + return nil, nil + } + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err == nil { + t.Fatal("expected an error when domain fails") + } + if !t1.systemDNSDirty { + t.Error("a failed revert must leave the resolver marked dirty") + } + if t1.systemDNSSet { + t.Error("systemDNSSet must stay false when the takeover did not complete") + } +} + +// A dirty resolver is retried before anything new is applied. +func TestConfigureSystemDNSRetriesDirtyBeforeApplying(t *testing.T) { + stubDNSRouting(t, nil) + calls := recorder(t, "") + + t1 := optedInTun() + t1.systemDNSDirty = true + + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil { + t.Fatalf("unexpected error: %v", err) + } + got := joined(*calls) + if !strings.HasPrefix(got, "resolvectl revert xray_tun") { + t.Errorf("expected the stale revert first, got %q", got) + } + if t1.systemDNSDirty { + t.Error("a successful retry should clear the dirty flag") + } + if !t1.systemDNSSet { + t.Error("the takeover should proceed once the retry succeeds") + } +} + +func TestUnsetSystemDNSReverts(t *testing.T) { + stubDNSRouting(t, nil) + calls := recorder(t, "") + + t1 := optedInTun() + if err := t1.ConfigureSystemDNS(context.Background(), "tun"); err != nil { + t.Fatalf("setup failed: %v", err) + } + *calls = nil + + t1.unsetSystemDNS() + if t1.systemDNSSet { + t.Error("systemDNSSet should be false after unset") + } + if got := joined(*calls); got != "resolvectl revert xray_tun" { + t.Errorf("unset calls = %q, want %q", got, "resolvectl revert xray_tun") + } + + t1.unsetSystemDNS() + if len(*calls) != 1 { + t.Errorf("unsetSystemDNS must be idempotent, got %q", joined(*calls)) + } +} + +func TestUnsetSystemDNSKeepsDirtyWhenRevertFails(t *testing.T) { + stubDNSRouting(t, nil) + calls := recorder(t, "revert") + + t1 := optedInTun() + t1.systemDNSSet = true + + t1.unsetSystemDNS() + if !t1.systemDNSDirty { + t.Error("a failed revert during unset must be remembered") + } + if got := joined(*calls); !strings.Contains(got, "resolvectl revert xray_tun") { + t.Errorf("expected a revert attempt, got %q", got) + } +} + +func TestSystemDNSAddrs(t *testing.T) { + tests := []struct { + name string + gateway []string + wantSource string + wantDNS string + wantOK bool + }{ + { + name: "ipv4 /30", + gateway: []string{"192.168.100.1/30"}, + wantSource: "192.168.100.1", + wantDNS: "192.168.100.2", + wantOK: true, + }, + { + name: "ipv4 /16", + gateway: []string{"10.0.0.1/16"}, + wantSource: "10.0.0.1", + wantDNS: "10.0.0.2", + wantOK: true, + }, + { + name: "first ipv4 wins", + gateway: []string{"fc00::1/64", "172.18.0.1/30"}, + wantSource: "172.18.0.1", + wantDNS: "172.18.0.2", + wantOK: true, + }, + { + name: "no gateway", + gateway: nil, + wantOK: false, + }, + { + name: "ipv6 only", + gateway: []string{"fc00::1/64"}, + wantSource: "fc00::1", + wantDNS: "fc00::2", + wantOK: true, + }, + { + name: "first ipv6 without ipv4", + gateway: []string{"fc00::1/64", "fd00::1/64"}, + wantSource: "fc00::1", + wantDNS: "fc00::2", + wantOK: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + source, dnsAddr, ok := systemDNSAddrs(tt.gateway) + if ok != tt.wantOK { + t.Fatalf("ok = %v, want %v", ok, tt.wantOK) + } + if !tt.wantOK { + return + } + if source.String() != tt.wantSource { + t.Errorf("source = %q, want %q", source.String(), tt.wantSource) + } + if dnsAddr.String() != tt.wantDNS { + t.Errorf("dns = %q, want %q", dnsAddr.String(), tt.wantDNS) + } + }) + } +} + +func TestBuildResolvectlArgs(t *testing.T) { + tests := []struct { + name string + action string + iface string + extra []string + want []string + }{ + { + name: "revert", + action: "revert", + iface: "xray_tun", + want: []string{"revert", "xray_tun"}, + }, + { + name: "dns single", + action: "dns", + iface: "xray_tun", + extra: []string{"192.168.100.2"}, + want: []string{"dns", "xray_tun", "192.168.100.2"}, + }, + { + name: "dns multiple", + action: "dns", + iface: "xray_tun", + extra: []string{"192.168.100.2", "fc00::2"}, + want: []string{"dns", "xray_tun", "192.168.100.2", "fc00::2"}, + }, + { + name: "domain wildcard", + action: "domain", + iface: "xray_tun", + extra: []string{"~."}, + want: []string{"domain", "xray_tun", "~."}, + }, + { + name: "default-route", + action: "default-route", + iface: "xray_tun", + extra: []string{"true"}, + want: []string{"default-route", "xray_tun", "true"}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := buildResolvectlArgs(tt.action, tt.iface, tt.extra...) + if len(got) != len(tt.want) { + t.Fatalf("args = %v, want %v", got, tt.want) + } + for i := range got { + if got[i] != tt.want[i] { + t.Errorf("args[%d] = %q, want %q (full: %v)", i, got[i], tt.want[i], got) + } + } + }) + } +} diff --git a/proxy/tun/tun_windows.go b/proxy/tun/tun_windows.go index 2097b6894..d0101c2f7 100644 --- a/proxy/tun/tun_windows.go +++ b/proxy/tun/tun_windows.go @@ -3,17 +3,25 @@ package tun import ( + "bytes" "context" "crypto/md5" "encoding/binary" go_errors "errors" "net" "net/netip" + "os/exec" + "path/filepath" + "slices" + "strconv" + "strings" "sync" + "syscall" "time" "unsafe" "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/transport/internet" "golang.org/x/sys/windows" "golang.zx2c4.com/wintun" "golang.zx2c4.com/wireguard/windows/tunnel/winipcfg" @@ -38,6 +46,10 @@ type WindowsTun struct { luid winipcfg.LUID cbr winipcfg.ChangeCallback cbi winipcfg.ChangeCallback + wfp windows.Handle + resolver *savedResolver + skipStop chan struct{} + skipDone chan struct{} closed bool } @@ -78,8 +90,13 @@ func open(name, desc string) (*wintun.Adapter, error) { // generate a deterministic GUID from the adapter name id := md5.Sum([]byte(name)) guid := (*windows.GUID)(unsafe.Pointer(&id[0])) + // try to open existing adapter by name + adapter, err := wintun.OpenAdapter(name) + if err == nil { + return adapter, nil + } // try to create adapter anew - adapter, err := wintun.CreateAdapter(name, desc, guid) + adapter, err = wintun.CreateAdapter(name, desc, guid) if err == nil { return adapter, nil } @@ -192,19 +209,105 @@ startOver: } } + // Windows lists the TUN's DNS servers among the system's ones, which Go's + // resolver queries for Xray's own lookups past the TUN, where they lead + // nowhere or back into Xray. Not skipped are those another interface uses + // as well, as that could leave no server at all. As those can change at + // any time, they are looked at again as often as Go rereads its servers. + if len(dns) > 0 { + skipped, err := tunOnlyDNS(t.luid, dns) + if err != nil { + skipped = dns + } + internet.SkipDNSServers(skipped) + t.skipStop, t.skipDone = make(chan struct{}), make(chan struct{}) + go func() { + defer close(t.skipDone) + ticker := time.NewTicker(5 * time.Second) + defer ticker.Stop() + for { + select { + case <-ticker.C: + if skipped, err := tunOnlyDNS(t.luid, dns); err == nil { + internet.SkipDNSServers(skipped) + } + case <-t.skipStop: + return + } + } + }() + } + + // Keep Windows from registering the TUN's addresses, and the host name + // with them, through dynamic DNS updates. Best effort. + if address4 || address6 { + if err := disableDNSRegistration(t.luid, dns); err != nil { + errors.LogDebugInner(context.Background(), err, "[tun] unable to disable DNS registration") + } + } + + // With autoSystemWfpBlockLeak, once the system routes lead to the TUN, + // keep DNS ("dns", if dns is set), and an IP version no route of which + // leads to the TUN ("misconfigtun"), from leaving through the other + // interfaces. Addresses do not matter: without one of a version in + // gateway, Windows gives the TUN a link-local one. + leaks := t.options.AutoSystemWfpBlockLeak + blockDNS := slices.Contains(leaks, "dns") && len(dns) > 0 + blockIPv4 := slices.Contains(leaks, "misconfigtun") && !route4 + blockIPv6 := slices.Contains(leaks, "misconfigtun") && !route6 + if (route4 || route6) && (blockDNS || blockIPv4 || blockIPv6) { + if t.wfp, err = blockLeaks(t.luid, blockDNS, blockIPv4, blockIPv6); err != nil { + var blocked []string + for _, b := range []struct { + on bool + what string + }{{blockDNS, "DNS"}, {blockIPv4, "IPv4"}, {blockIPv6, "IPv6"}} { + if b.on { + blocked = append(blocked, b.what) + } + } + // Rather no TUN than a leaking one. + return errors.New("unable to block ", strings.Join(blocked, " and "), " outside the TUN (remove autoSystemWfpBlockLeak to run without)").Base(err) + } + errors.LogInfo(context.Background(), "[tun] outside the TUN, blocked DNS: ", blockDNS, ", blocked IPv4: ", blockIPv4, ", blocked IPv6: ", blockIPv6) + if blockDNS { + covered := slices.Clone(addresses) + for _, route := range routesData { + covered = append(covered, route.Destination) + } + for _, server := range dnsOutsideTUN(dns, covered) { + errors.LogWarning(context.Background(), "[tun] DNS server ", server, " is in neither gateway nor autoSystemRoutingTable, so queries to it cannot go through the TUN and are blocked") + } + // With updater, the dialer controllers bind Xray's own sockets + // to the physical interface. + if updater != nil { + t.resolver = resolveOnOwn() + } + } + } + if len(dns) > 0 || route4 || route6 { + if err := flushDNSCache(); err != nil { + errors.LogInfoInner(context.Background(), err, "[tun] unable to flush DNS cache") + } + } + if updater != nil { - t.cbr, err = winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) { + // Only a registered callback goes into the fields: a nil pointer in + // them would not compare equal to nil in Close. + cbr, err := winipcfg.RegisterRouteChangeCallback(func(notificationType winipcfg.MibNotificationType, route *winipcfg.MibIPforwardRow2) { updater.Update() }) if err != nil { return err } - t.cbi, err = winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) { + t.cbr = cbr + cbi, err := winipcfg.RegisterInterfaceChangeCallback(func(notificationType winipcfg.MibNotificationType, iface *winipcfg.MibIPInterfaceRow) { updater.Update() }) if err != nil { return err } + t.cbi = cbi } return nil } @@ -231,6 +334,20 @@ func (t *WindowsTun) Close() error { t.luid.FlushIPAddresses(windows.AF_INET6) t.luid.FlushDNS(windows.AF_INET6) } + if t.wfp != 0 { + closeWFPEngine(t.wfp) + } + if t.resolver != nil { + t.resolver.restore() + } + if t.skipStop != nil { + close(t.skipStop) + <-t.skipDone + } + internet.SkipDNSServers(nil) + if len(t.options.DNS) > 0 || len(t.options.AutoSystemRoutingTable) > 0 { + flushDNSCache() + } if t.session != (wintun.Session{}) { t.session.End() } @@ -240,6 +357,121 @@ func (t *WindowsTun) Close() error { return nil } +type savedResolver struct { + preferGo bool + dial func(ctx context.Context, network, address string) (net.Conn, error) +} + +// resolveOnOwn has Go resolve the names Xray would otherwise ask Windows for, +// on Xray's own sockets, which the dialer controllers bind to the physical +// interface, and skipping the TUN's DNS servers, as localdns does. Windows' +// resolver runs in the DNS Client service, whose queries the DNS filter lets +// through the TUN only, so Xray's own lookups, like of an outbound's server +// domain, would go into Xray again and could end up waiting on themselves. +// +// It changes net.DefaultResolver for the whole process, which covers every +// lookup that would reach Windows' resolver; restore undoes it. +func resolveOnOwn() *savedResolver { + saved := &savedResolver{net.DefaultResolver.PreferGo, net.DefaultResolver.Dial} + dialer := &net.Dialer{Control: func(network, address string, c syscall.RawConn) error { + for _, ctl := range internet.Controllers { + if err := ctl(network, address, c); err != nil { + return err + } + } + return nil + }} + // Go's resolver moves on to the next server right away when a dial fails. + net.DefaultResolver.Dial = func(ctx context.Context, network, address string) (net.Conn, error) { + if internet.IsSkippedDNSServer(address) { + return nil, errors.New("skipped DNS server ", address) + } + return dialer.DialContext(ctx, network, address) + } + net.DefaultResolver.PreferGo = true + return saved +} + +func (s *savedResolver) restore() { + net.DefaultResolver.PreferGo = s.preferGo + net.DefaultResolver.Dial = s.dial +} + +// tunOnlyDNS returns those of servers, the TUN's DNS servers, that Go's +// resolver does not also get from another interface: one that is up and has +// a gateway, as it reads them. +func tunOnlyDNS(tun winipcfg.LUID, servers []netip.Addr) ([]netip.Addr, error) { + adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways) + if err != nil { + return nil, err + } + var others []netip.Addr + for _, adapter := range adapters { + if adapter.LUID == tun || adapter.OperStatus != winipcfg.IfOperStatusUp || adapter.FirstGatewayAddress == nil { + continue + } + for server := adapter.FirstDNSServerAddress; server != nil; server = server.Next { + if addr, ok := netip.AddrFromSlice(server.Address.IP()); ok { + others = append(others, addr.Unmap()) + } + } + } + return slices.DeleteFunc(slices.Clone(servers), func(server netip.Addr) bool { + return slices.Contains(others, server.Unmap()) + }), nil +} + +// disableDNSRegistration turns off the dynamic DNS registration of the +// interface's addresses. dns are its DNS servers. +func disableDNSRegistration(luid winipcfg.LUID, dns []netip.Addr) error { + guid, err := luid.GUID() + if err != nil { + return err + } + err = winipcfg.SetInterfaceDnsSettings(*guid, &winipcfg.DnsInterfaceSettings{ + Version: winipcfg.DnsInterfaceSettingsVersion1, + Flags: winipcfg.DnsInterfaceSettingsFlagRegistrationEnabled, + }) + if err == nil || !go_errors.Is(err, windows.ERROR_PROC_NOT_FOUND) { + return err + } + return disableDNSRegistrationByNetsh(luid, dns) +} + +// disableDNSRegistrationByNetsh does it for Windows before 10 1809, which +// lacks SetInterfaceDnsSettings. The setting is the interface's, not the +// address family's, but netsh only applies it along with a DNS server, which +// replaces the IPv4 ones, so they are set again afterwards. +func disableDNSRegistrationByNetsh(luid winipcfg.LUID, dns []netip.Addr) error { + row, err := luid.Interface() + if err != nil { + return err + } + server := "127.0.0.1" // any will do when there is no IPv4 one + if i := slices.IndexFunc(dns, netip.Addr.Is4); i >= 0 { + server = dns[i].String() + } + err = runNetsh("interface", "ipv4", "set", "dnsservers", "name="+strconv.FormatUint(uint64(row.InterfaceIndex), 10), "source=static", "address="+server, "register=none", "validate=no") + return errors.Combine(err, luid.SetDNS(windows.AF_INET, dns, nil)) +} + +// runNetsh runs netsh.exe from the system directory. netsh reports some +// failures, like a syntax error, only in its output, even with exit code 0, +// so any output counts as a failure. +func runNetsh(args ...string) error { + system32, err := windows.GetSystemDirectory() + if err != nil { + return err + } + cmd := exec.Command(filepath.Join(system32, "netsh.exe"), args...) + cmd.SysProcAttr = &syscall.SysProcAttr{HideWindow: true} + output, err := cmd.CombinedOutput() + if output = bytes.TrimSpace(output); err != nil || len(output) > 0 { + return errors.New("netsh ", strings.Join(args, " "), ": ", string(output)).Base(err) + } + return nil +} + func (t *WindowsTun) Name() (string, error) { row, err := t.luid.Interface() if err != nil { diff --git a/proxy/tun/tun_windows_wfp.go b/proxy/tun/tun_windows_wfp.go new file mode 100644 index 000000000..4cf952b41 --- /dev/null +++ b/proxy/tun/tun_windows_wfp.go @@ -0,0 +1,471 @@ +//go:build windows + +package tun + +import ( + "net/netip" + "os" + "runtime" + "slices" + "unsafe" + + "github.com/xtls/xray-core/common/errors" + "golang.org/x/sys/windows" + "golang.zx2c4.com/wireguard/windows/tunnel/winipcfg" +) + +var ( + modfwpuclnt = windows.NewLazySystemDLL("fwpuclnt.dll") + moddnsapi = windows.NewLazySystemDLL("dnsapi.dll") + + procFwpmEngineOpen0 = modfwpuclnt.NewProc("FwpmEngineOpen0") + procFwpmEngineClose0 = modfwpuclnt.NewProc("FwpmEngineClose0") + procFwpmTransactionBegin0 = modfwpuclnt.NewProc("FwpmTransactionBegin0") + procFwpmTransactionCommit0 = modfwpuclnt.NewProc("FwpmTransactionCommit0") + procFwpmTransactionAbort0 = modfwpuclnt.NewProc("FwpmTransactionAbort0") + procFwpmSubLayerAdd0 = modfwpuclnt.NewProc("FwpmSubLayerAdd0") + procFwpmFilterAdd0 = modfwpuclnt.NewProc("FwpmFilterAdd0") + procFwpmGetAppIdFromFileName0 = modfwpuclnt.NewProc("FwpmGetAppIdFromFileName0") + procFwpmFreeMemory0 = modfwpuclnt.NewProc("FwpmFreeMemory0") + procDnsFlushResolverCache = moddnsapi.NewProc("DnsFlushResolverCache") +) + +// fwptypes.h and fwpmtypes.h +const ( + rpcCAuthnWinNT = 10 // RPC_C_AUTHN_WINNT + fwpmSessionFlagDynamic = 1 // FWPM_SESSION_FLAG_DYNAMIC + fwpmFilterFlagClearActionRight = 8 // FWPM_FILTER_FLAG_CLEAR_ACTION_RIGHT + + fwpUint8 = 1 // FWP_UINT8 + fwpUint16 = 2 // FWP_UINT16 + fwpUint32 = 3 // FWP_UINT32 + fwpUint64 = 4 // FWP_UINT64 + fwpByteArray16Type = 11 // FWP_BYTE_ARRAY16_TYPE + fwpByteBlobType = 12 // FWP_BYTE_BLOB_TYPE + fwpSecurityDescriptorType = 14 // FWP_SECURITY_DESCRIPTOR_TYPE + + fwpMatchEqual = 0 // FWP_MATCH_EQUAL + fwpMatchFlagsAllSet = 6 // FWP_MATCH_FLAGS_ALL_SET + + fwpConditionFlagIsLoopback = 1 // FWP_CONDITION_FLAG_IS_LOOPBACK + + fwpActionBlock = 0x1001 // FWP_ACTION_BLOCK + fwpActionPermit = 0x1002 // FWP_ACTION_PERMIT +) + +// fwpmu.h +var ( + fwpmLayerALEAuthConnectV4 = windows.GUID{Data1: 0xc38d57d1, Data2: 0x05a7, Data3: 0x4c33, Data4: [8]byte{0x90, 0x4f, 0x7f, 0xbc, 0xee, 0xe6, 0x0e, 0x82}} + fwpmLayerALEAuthConnectV6 = windows.GUID{Data1: 0x4a72393b, Data2: 0x319f, Data3: 0x44bc, Data4: [8]byte{0x84, 0xc3, 0xba, 0x54, 0xdc, 0xb3, 0xb6, 0xb4}} + fwpmLayerALEAuthRecvAcceptV4 = windows.GUID{Data1: 0xe1cd9fe7, Data2: 0xf4b5, Data3: 0x4273, Data4: [8]byte{0x96, 0xc0, 0x59, 0x2e, 0x48, 0x7b, 0x86, 0x50}} + fwpmLayerALEAuthRecvAcceptV6 = windows.GUID{Data1: 0xa3b42c97, Data2: 0x9f04, Data3: 0x4672, Data4: [8]byte{0xb8, 0x7e, 0xce, 0xe9, 0xc4, 0x83, 0x25, 0x7f}} + + fwpmConditionFlags = windows.GUID{Data1: 0x632ce23b, Data2: 0x5167, Data3: 0x435c, Data4: [8]byte{0x86, 0xd7, 0xe9, 0x03, 0x68, 0x4a, 0xa8, 0x0c}} + fwpmConditionIPArrivalInterface = windows.GUID{Data1: 0x618a9b6d, Data2: 0x386b, Data3: 0x4136, Data4: [8]byte{0xad, 0x6e, 0xb5, 0x15, 0x87, 0xcf, 0xb1, 0xcd}} + fwpmConditionIPLocalInterface = windows.GUID{Data1: 0x4cd62a49, Data2: 0x59c3, Data3: 0x4969, Data4: [8]byte{0xb7, 0xf3, 0xbd, 0xa5, 0xd3, 0x28, 0x90, 0xa4}} + fwpmConditionIPLocalPort = windows.GUID{Data1: 0x0c1ba1af, Data2: 0x5765, Data3: 0x453f, Data4: [8]byte{0xaf, 0x22, 0xa8, 0xf7, 0x91, 0xac, 0x77, 0x5b}} // also FWPM_CONDITION_ICMP_TYPE + fwpmConditionIPNexthopInterface = windows.GUID{Data1: 0x93ae8f5b, Data2: 0x7f6f, Data3: 0x4719, Data4: [8]byte{0x98, 0xc8, 0x14, 0xe9, 0x74, 0x29, 0xef, 0x04}} + fwpmConditionIPProtocol = windows.GUID{Data1: 0x3971ef2b, Data2: 0x623e, Data3: 0x4f9a, Data4: [8]byte{0x8c, 0xb1, 0x6e, 0x79, 0xb8, 0x06, 0xb9, 0xa7}} + fwpmConditionIPRemoteAddress = windows.GUID{Data1: 0xb235ae9a, Data2: 0x1d64, Data3: 0x49b8, Data4: [8]byte{0xa4, 0x4c, 0x5f, 0xf3, 0xd9, 0x09, 0x50, 0x45}} + fwpmConditionIPRemotePort = windows.GUID{Data1: 0xc35a604d, Data2: 0xd22b, Data3: 0x4e1a, Data4: [8]byte{0x91, 0xb4, 0x68, 0xf6, 0x74, 0xee, 0x67, 0x4b}} // also FWPM_CONDITION_ICMP_CODE + fwpmConditionALEAppID = windows.GUID{Data1: 0xd78e1e87, Data2: 0x8644, Data3: 0x4ea5, Data4: [8]byte{0x94, 0x37, 0xd8, 0x09, 0xec, 0xef, 0xc9, 0x71}} + fwpmConditionALEUserID = windows.GUID{Data1: 0xaf043a0a, Data2: 0xb34d, Data3: 0x4f86, Data4: [8]byte{0x97, 0x9c, 0xc9, 0x03, 0x71, 0xaf, 0x6e, 0x66}} +) + +// dnsClientSID is the SID of Windows' DNS Client service, NT SERVICE\Dnscache. +// Service SIDs derive from the service name, so it is the same everywhere (sc +// showsid dnscache). +const dnsClientSID = "S-1-5-80-859482183-879914841-863379149-1145462774-2388618682" + +// ff02::1:2, where DHCPv6 clients send to. A package-level variable never +// moves, so conditions may refer to it through uintptr. +var ipv6AllDHCPv6Servers = [16]byte{0xff, 0x02, 13: 0x01, 15: 0x02} + +type fwpByteBlob struct { + size uint32 + data *byte +} + +// fwpValue0 is FWP_VALUE0 as well as FWP_CONDITION_VALUE0. Their union holds +// a scalar of at most 32 bits, or a pointer for the larger types. +type fwpValue0 struct { + typ uint32 + value uintptr +} + +type fwpmDisplayData0 struct { + name *uint16 + description *uint16 +} + +type fwpmSession0 struct { + sessionKey windows.GUID + displayData fwpmDisplayData0 + flags uint32 + txnWaitTimeoutInMSec uint32 + processID uint32 + sid *windows.SID + username *uint16 + kernelMode int32 +} + +type fwpmSublayer0 struct { + subLayerKey windows.GUID + displayData fwpmDisplayData0 + flags uint32 + providerKey *windows.GUID + providerData fwpByteBlob + weight uint16 +} + +type fwpmFilterCondition0 struct { + fieldKey windows.GUID + matchType uint32 + conditionValue fwpValue0 +} + +type fwpmAction0 struct { + typ uint32 + filterType windows.GUID +} + +type fwpmFilter0 struct { + filterKey windows.GUID + displayData fwpmDisplayData0 + flags uint32 + providerKey *windows.GUID + providerData fwpByteBlob + layerKey windows.GUID + subLayerKey windows.GUID + weight fwpValue0 + numFilterConditions uint32 + filterCondition *fwpmFilterCondition0 + action fwpmAction0 + _ uint32 // C aligns the following union to 8 bytes, as it holds a UINT64 + providerContextKey windows.GUID + reserved *windows.GUID + _ [8 - unsafe.Sizeof(uintptr(0))]byte // and filterId as well, also on 32-bit + filterID uint64 + effectiveWeight fwpValue0 +} + +// fwpmResult converts the DWORD status the Fwpm functions return. +func fwpmResult(r1, _ uintptr, _ error) error { + if r1 != 0 { + return windows.Errno(r1) + } + return nil +} + +func utf16Ptr(s string) *uint16 { + p, _ := windows.UTF16PtrFromString(s) + return p +} + +func condition(field *windows.GUID, typ uint32, value uintptr) fwpmFilterCondition0 { + return fwpmFilterCondition0{ + fieldKey: *field, + matchType: fwpMatchEqual, + conditionValue: fwpValue0{typ: typ, value: value}, + } +} + +// blockLeaks keeps traffic from leaving through interfaces other than tun, +// for every program but Xray itself, whose outbounds (DNS included) use the +// other interfaces on purpose: +// +// - dns: DNS (port 53) may only go through the TUN. Windows sends a name +// query to the DNS servers of all interfaces, not only to those of the TUN: +// to the first server of each interface, then to all of them when no answer +// arrives within a second or two. It sends the queries for the servers of +// an interface out through that interface, whatever the routes say, and +// other programs reach an on-link resolver, like 192.168.1.1 from DHCP, +// through its LAN route, which is more specific than the TUN's default +// route. Since Windows 11 and Server 2022, Windows may also send its +// queries over HTTPS or TLS, so there its DNS Client service may not +// connect outside the TUN at all, except for name resolution on the local +// link (mDNS, LLMNR). +// - ipv4, ipv6: no IPv4, or no IPv6, at all, in either direction, for a TUN +// that no route of it leads to, except loopback and what Windows itself +// needs on the local link (DHCP, and for IPv6 neighbor and multicast +// listener discovery), none of which can leave it. The TUN carries what +// is routed to it even without an address of that IP version in gateway: +// Windows gives it link-local ones itself, an IPv6 one at once, an IPv4 +// one from 169.254.0.0/16 after some seconds (until then, IPv4 routed to +// the TUN is unreachable). +// +// The filters live in a dynamic WFP session: closing the returned engine handle +// with closeWFPEngine deletes them, and so does Windows when the process dies. +func blockLeaks(tun winipcfg.LUID, dns, ipv4, ipv6 bool) (windows.Handle, error) { + engine, err := openWFPEngine() + if err != nil { + return 0, err + } + if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil { + closeWFPEngine(engine) + return 0, errors.New("FwpmTransactionBegin0 failed").Base(err) + } + err = addLeakFilters(engine, tun, dns, ipv4, ipv6) + if err == nil { + if err = fwpmResult(procFwpmTransactionCommit0.Call(uintptr(engine))); err != nil { + err = errors.New("FwpmTransactionCommit0 failed").Base(err) + } + } + if err != nil { + procFwpmTransactionAbort0.Call(uintptr(engine)) + closeWFPEngine(engine) + return 0, err + } + return engine, nil +} + +func openWFPEngine() (windows.Handle, error) { + if err := modfwpuclnt.Load(); err != nil { + return 0, err + } + // txnWaitTimeoutInMSec stays 0 for BFE's default, so that a transaction + // held by another program cannot hang the start forever. + session := fwpmSession0{ + displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")}, + flags: fwpmSessionFlagDynamic, + } + var engine windows.Handle + if err := fwpmResult(procFwpmEngineOpen0.Call(0, rpcCAuthnWinNT, 0, uintptr(unsafe.Pointer(&session)), uintptr(unsafe.Pointer(&engine)))); err != nil { + return 0, errors.New("FwpmEngineOpen0 failed").Base(err) + } + return engine, nil +} + +func closeWFPEngine(engine windows.Handle) { + procFwpmEngineClose0.Call(uintptr(engine)) +} + +// addLeakFilters adds the filters of blockLeaks in a sublayer of their own. +// blockLeaks runs it in a transaction, so that they take effect all at once. +func addLeakFilters(engine windows.Handle, tun winipcfg.LUID, dns, ipv4, ipv6 bool) error { + exe, err := os.Executable() + if err != nil { + return err + } + exePath, err := windows.UTF16PtrFromString(exe) + if err != nil { + return err + } + var appID *fwpByteBlob + if err := fwpmResult(procFwpmGetAppIdFromFileName0.Call(uintptr(unsafe.Pointer(exePath)), uintptr(unsafe.Pointer(&appID)))); err != nil { + return errors.New("FwpmGetAppIdFromFileName0 failed for ", exe).Base(err) + } + defer func() { procFwpmFreeMemory0.Call(uintptr(unsafe.Pointer(&appID))) }() + + sublayer := fwpmSublayer0{ + displayData: fwpmDisplayData0{name: utf16Ptr("Xray TUN")}, + weight: 0xffff, + } + if sublayer.subLayerKey, err = windows.GenerateGUID(); err != nil { + return err + } + if err := fwpmResult(procFwpmSubLayerAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&sublayer)), 0)); err != nil { + return errors.New("FwpmSubLayerAdd0 failed").Base(err) + } + add := func(layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error { + return addFilter(engine, &sublayer.subLayerKey, layer, "Xray TUN: "+name, flags, action, weight, conditions...) + } + + var pinner runtime.Pinner + defer pinner.Unpin() + tunLUID := new(uint64) + *tunLUID = uint64(tun) + pinner.Pin(tunLUID) // the condition only holds it as uintptr + + // The heaviest matching filter of a sublayer decides. All sublayers have + // their say, though, and a block in any of them beats a permit, unless + // the permit is hard: it clears the action right, and then the blocks of + // lower sublayers, Windows Firewall rules among them, no longer override + // it, only a callout's veto does. Xray's own connections out get such a + // hard permit. Connections from outside to Xray get an ordinary one, so + // that firewalls keep guarding its inbounds. + self := condition(&fwpmConditionALEAppID, fwpByteBlobType, uintptr(unsafe.Pointer(appID))) + dns53 := condition(&fwpmConditionIPRemotePort, fwpUint16, 53) + // DNS goes through the TUN when its local address is the TUN's, and it + // also leaves, or arrives, through the TUN. The local address alone + // decides by default, but with weak host sending or receiving enabled, + // packets of the TUN's address can use other interfaces. (The next hop, + // the interface replies would leave by, is not known for arriving ones.) + onTUN := func(field *windows.GUID) fwpmFilterCondition0 { + return condition(field, fwpUint64, uintptr(unsafe.Pointer(tunLUID))) + } + out := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)} + in := []fwpmFilterCondition0{dns53, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPArrivalInterface)} + for _, layer := range []struct { + key *windows.GUID + selfFlags uint32 + throughTUN []fwpmFilterCondition0 + }{ + {&fwpmLayerALEAuthConnectV4, fwpmFilterFlagClearActionRight, out}, + {&fwpmLayerALEAuthRecvAcceptV4, 0, in}, + {&fwpmLayerALEAuthConnectV6, fwpmFilterFlagClearActionRight, out}, + {&fwpmLayerALEAuthRecvAcceptV6, 0, in}, + } { + if err := add(layer.key, "permit Xray", layer.selfFlags, fwpActionPermit, 4, self); err != nil { + return err + } + if dns { + if err := add(layer.key, "permit DNS through the TUN", 0, fwpActionPermit, 3, layer.throughTUN...); err != nil { + return err + } + if err := add(layer.key, "block DNS", 0, fwpActionBlock, 2, dns53); err != nil { + return err + } + } + } + + // Since Windows 11 and Server 2022 (build 20348), the DNS Client service + // may also send the queries for an interface's servers over HTTPS or TLS, + // out through that interface and to any port. So there it may only + // connect through the TUN, except for mDNS and LLMNR, which stay on the + // local link (over an IP version only while it is not blocked altogether). + // Earlier versions only query port 53, and may run the service in one + // process with others, which the filters would catch as well. Like + // Windows Firewall's rules for it, they recognize the service by its SID, + // which Windows puts in the token of its process: the security descriptor + // grants that SID the right to match (FWP_ACTRL_MATCH_FILTER, CC in SDDL). + if _, _, build := windows.RtlGetNtVersionNumbers(); dns && build >= 20348 { + sd, err := windows.SecurityDescriptorFromString("O:SYG:SYD:(A;;CCRC;;;" + dnsClientSID + ")") + if err != nil { + return err + } + sdBlob := &fwpByteBlob{size: sd.Length(), data: (*byte)(unsafe.Pointer(sd))} + pinner.Pin(sdBlob) // the condition only holds it as uintptr + dnsClient := condition(&fwpmConditionALEUserID, fwpSecurityDescriptorType, uintptr(unsafe.Pointer(sdBlob))) + // Conditions on the same field match when any of them does. + mdnsLLMNR := []fwpmFilterCondition0{dnsClient, condition(&fwpmConditionIPRemotePort, fwpUint16, 5353), condition(&fwpmConditionIPRemotePort, fwpUint16, 5355)} + for _, layer := range []struct { + key *windows.GUID + localLink bool + }{ + {&fwpmLayerALEAuthConnectV4, !ipv4}, + {&fwpmLayerALEAuthConnectV6, !ipv6}, + } { + if err := add(layer.key, "permit the DNS Client service through the TUN", 0, fwpActionPermit, 3, dnsClient, onTUN(&fwpmConditionIPLocalInterface), onTUN(&fwpmConditionIPNexthopInterface)); err != nil { + return err + } + if layer.localLink { + if err := add(layer.key, "permit the DNS Client service's mDNS and LLMNR", 0, fwpActionPermit, 3, mdnsLLMNR...); err != nil { + return err + } + } + if err := add(layer.key, "block the DNS Client service", 0, fwpActionBlock, 2, dnsClient); err != nil { + return err + } + } + } + + // Both directions: replies to a connection accepted from outside would + // leave through the physical link as well. + loopback := fwpmFilterCondition0{ + fieldKey: fwpmConditionFlags, + matchType: fwpMatchFlagsAllSet, + conditionValue: fwpValue0{typ: fwpUint32, value: fwpConditionFlagIsLoopback}, + } + if ipv4 { + // DHCP keeps the addresses of the other interfaces, which Xray's own + // connections use. + dhcp := []fwpmFilterCondition0{ + condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP), + condition(&fwpmConditionIPLocalPort, fwpUint16, 68), + condition(&fwpmConditionIPRemotePort, fwpUint16, 67), + } + for _, layer := range []*windows.GUID{&fwpmLayerALEAuthConnectV4, &fwpmLayerALEAuthRecvAcceptV4} { + if err := add(layer, "permit IPv4 loopback", 0, fwpActionPermit, 1, loopback); err != nil { + return err + } + if err := add(layer, "permit DHCP", 0, fwpActionPermit, 1, dhcp...); err != nil { + return err + } + if err := add(layer, "block IPv4", 0, fwpActionBlock, 0); err != nil { + return err + } + } + } + if ipv6 { + // Neighbor and multicast listener discovery, ICMPv6 130-137 and 143, + // whose type and code sit where the local and remote port are. + discovery := []fwpmFilterCondition0{condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_ICMPV6)} + for _, typ := range []uintptr{130, 131, 132, 133, 134, 135, 136, 137, 143} { + discovery = append(discovery, condition(&fwpmConditionIPLocalPort, fwpUint16, typ)) + } + discovery = append(discovery, condition(&fwpmConditionIPRemotePort, fwpUint16, 0)) + dhcpv6 := []fwpmFilterCondition0{ + condition(&fwpmConditionIPProtocol, fwpUint8, windows.IPPROTO_UDP), + condition(&fwpmConditionIPLocalPort, fwpUint16, 546), + condition(&fwpmConditionIPRemotePort, fwpUint16, 547), + } + for _, direction := range []struct { + layer *windows.GUID + dhcpv6 []fwpmFilterCondition0 + }{ + // The client sends to the servers' multicast address, and they + // answer from their own. + {&fwpmLayerALEAuthConnectV6, slices.Concat(dhcpv6, []fwpmFilterCondition0{condition(&fwpmConditionIPRemoteAddress, fwpByteArray16Type, uintptr(unsafe.Pointer(&ipv6AllDHCPv6Servers)))})}, + {&fwpmLayerALEAuthRecvAcceptV6, dhcpv6}, + } { + if err := add(direction.layer, "permit IPv6 loopback", 0, fwpActionPermit, 1, loopback); err != nil { + return err + } + if err := add(direction.layer, "permit IPv6 neighbor and multicast listener discovery", 0, fwpActionPermit, 1, discovery...); err != nil { + return err + } + if err := add(direction.layer, "permit DHCPv6", 0, fwpActionPermit, 1, direction.dhcpv6...); err != nil { + return err + } + if err := add(direction.layer, "block IPv6", 0, fwpActionBlock, 0); err != nil { + return err + } + } + } + return nil +} + +func addFilter(engine windows.Handle, sublayer, layer *windows.GUID, name string, flags, action uint32, weight uint8, conditions ...fwpmFilterCondition0) error { + filter := fwpmFilter0{ + displayData: fwpmDisplayData0{name: utf16Ptr(name)}, + flags: flags, + layerKey: *layer, + subLayerKey: *sublayer, + weight: fwpValue0{typ: fwpUint8, value: uintptr(weight)}, + numFilterConditions: uint32(len(conditions)), + action: fwpmAction0{typ: action}, + } + if len(conditions) > 0 { + filter.filterCondition = &conditions[0] + } + if err := fwpmResult(procFwpmFilterAdd0.Call(uintptr(engine), uintptr(unsafe.Pointer(&filter)), 0, 0)); err != nil { + return errors.New("FwpmFilterAdd0 failed for ", name).Base(err) + } + return nil +} + +// dnsOutsideTUN returns the servers outside all of prefixes, the TUN's own +// subnets and routes: queries to them cannot go through the TUN. +func dnsOutsideTUN(servers []netip.Addr, prefixes []netip.Prefix) []netip.Addr { + var outside []netip.Addr + for _, server := range servers { + server = server.Unmap() + if !slices.ContainsFunc(prefixes, func(p netip.Prefix) bool { return p.Contains(server) }) { + outside = append(outside, server) + } + } + return outside +} + +// flushDNSCache drops the answers Windows cached so far, like ipconfig +// /flushdns, so that names get resolved again with the current DNS setup. +func flushDNSCache() error { + if err := procDnsFlushResolverCache.Find(); err != nil { + return err + } + if r, _, err := procDnsFlushResolverCache.Call(); r == 0 { + return err + } + return nil +} diff --git a/proxy/tun/tun_windows_wfp_test.go b/proxy/tun/tun_windows_wfp_test.go new file mode 100644 index 000000000..57fc005fe --- /dev/null +++ b/proxy/tun/tun_windows_wfp_test.go @@ -0,0 +1,206 @@ +//go:build windows + +package tun + +import ( + "context" + go_errors "errors" + "net" + "net/netip" + "slices" + "testing" + "unsafe" + + "github.com/xtls/xray-core/transport/internet" + "golang.org/x/sys/windows" + "golang.zx2c4.com/wireguard/windows/tunnel/winipcfg" +) + +// The WFP structures are handed to fwpuclnt.dll as they are, so their layout +// has to match what MSVC produces for 64-bit and for 32-bit Windows. +func TestWFPStructLayout(t *testing.T) { + check := func(name string, got, want64, want32 []uintptr) { + t.Helper() + want := want32 + if unsafe.Sizeof(uintptr(0)) == 8 { + want = want64 + } + if !slices.Equal(got, want) { + t.Errorf("%s: size and offsets are %v, want %v", name, got, want) + } + } + + var blob fwpByteBlob + check("FWP_BYTE_BLOB", + []uintptr{unsafe.Sizeof(blob), unsafe.Offsetof(blob.data)}, + []uintptr{16, 8}, []uintptr{8, 4}) + + var value fwpValue0 + check("FWP_VALUE0", + []uintptr{unsafe.Sizeof(value), unsafe.Offsetof(value.value)}, + []uintptr{16, 8}, []uintptr{8, 4}) + + var display fwpmDisplayData0 + check("FWPM_DISPLAY_DATA0", + []uintptr{unsafe.Sizeof(display), unsafe.Offsetof(display.description)}, + []uintptr{16, 8}, []uintptr{8, 4}) + + var action fwpmAction0 + check("FWPM_ACTION0", + []uintptr{unsafe.Sizeof(action), unsafe.Offsetof(action.filterType)}, + []uintptr{20, 4}, []uintptr{20, 4}) + + var cond fwpmFilterCondition0 + check("FWPM_FILTER_CONDITION0", + []uintptr{unsafe.Sizeof(cond), unsafe.Offsetof(cond.matchType), unsafe.Offsetof(cond.conditionValue)}, + []uintptr{40, 16, 24}, []uintptr{28, 16, 20}) + + var session fwpmSession0 + check("FWPM_SESSION0", + []uintptr{ + unsafe.Sizeof(session), unsafe.Offsetof(session.displayData), unsafe.Offsetof(session.flags), + unsafe.Offsetof(session.txnWaitTimeoutInMSec), unsafe.Offsetof(session.processID), unsafe.Offsetof(session.sid), + unsafe.Offsetof(session.username), unsafe.Offsetof(session.kernelMode), + }, + []uintptr{72, 16, 32, 36, 40, 48, 56, 64}, + []uintptr{48, 16, 24, 28, 32, 36, 40, 44}) + + var sublayer fwpmSublayer0 + check("FWPM_SUBLAYER0", + []uintptr{ + unsafe.Sizeof(sublayer), unsafe.Offsetof(sublayer.displayData), unsafe.Offsetof(sublayer.flags), + unsafe.Offsetof(sublayer.providerKey), unsafe.Offsetof(sublayer.providerData), unsafe.Offsetof(sublayer.weight), + }, + []uintptr{72, 16, 32, 40, 48, 64}, + []uintptr{44, 16, 24, 28, 32, 40}) + + var filter fwpmFilter0 + check("FWPM_FILTER0", + []uintptr{ + unsafe.Sizeof(filter), unsafe.Offsetof(filter.displayData), unsafe.Offsetof(filter.flags), + unsafe.Offsetof(filter.providerKey), unsafe.Offsetof(filter.providerData), unsafe.Offsetof(filter.layerKey), + unsafe.Offsetof(filter.subLayerKey), unsafe.Offsetof(filter.weight), unsafe.Offsetof(filter.numFilterConditions), + unsafe.Offsetof(filter.filterCondition), unsafe.Offsetof(filter.action), unsafe.Offsetof(filter.providerContextKey), + unsafe.Offsetof(filter.reserved), unsafe.Offsetof(filter.filterID), unsafe.Offsetof(filter.effectiveWeight), + }, + []uintptr{200, 16, 32, 40, 48, 64, 80, 96, 112, 120, 128, 152, 168, 176, 184}, + []uintptr{152, 16, 24, 28, 32, 40, 56, 72, 80, 84, 88, 112, 128, 136, 144}) +} + +// TestLeakFiltersAccepted has WFP validate the filters by adding them inside a +// transaction that is then aborted, which leaves the system untouched. Adding +// filters requires an elevated process. +func TestLeakFiltersAccepted(t *testing.T) { + skipUnlessElevated := func(err error) { + t.Helper() + if go_errors.Is(err, windows.ERROR_ACCESS_DENIED) { + t.Skipf("WFP filters can only be added by an elevated process: %v", err) + } + t.Fatal(err) + } + + engine, err := openWFPEngine() + if err != nil { + skipUnlessElevated(err) + } + defer closeWFPEngine(engine) + if err := fwpmResult(procFwpmTransactionBegin0.Call(uintptr(engine), 0)); err != nil { + skipUnlessElevated(err) + } + defer procFwpmTransactionAbort0.Call(uintptr(engine)) + + // Any interface stands in for the TUN; the loopback one always exists. + loopback, err := winipcfg.LUIDFromIndex(1) + if err != nil { + t.Fatal(err) + } + if err := addLeakFilters(engine, loopback, true, true, true); err != nil { + skipUnlessElevated(err) + } +} + +func TestDNSClientSID(t *testing.T) { + sid, _, _, err := windows.LookupSID("", `NT SERVICE\Dnscache`) + if err != nil { + t.Fatal(err) + } + if sid.String() != dnsClientSID { + t.Errorf(`NT SERVICE\Dnscache is %v, not %v`, sid, dnsClientSID) + } +} + +func TestDNSOutsideTUN(t *testing.T) { + prefixes := []netip.Prefix{ + netip.MustParsePrefix("198.51.100.1/30"), // gateway, not masked + netip.MustParsePrefix("203.0.113.0/24"), // route + } + servers := []netip.Addr{ + netip.MustParseAddr("198.51.100.2"), + netip.MustParseAddr("203.0.113.53"), + netip.MustParseAddr("::ffff:203.0.113.54"), + netip.MustParseAddr("8.8.8.8"), + netip.MustParseAddr("2001:db8::53"), + } + want := []netip.Addr{netip.MustParseAddr("8.8.8.8"), netip.MustParseAddr("2001:db8::53")} + if got := dnsOutsideTUN(servers, prefixes); !slices.Equal(got, want) { + t.Errorf("got %v, want %v", got, want) + } +} + +func TestResolveOnOwn(t *testing.T) { + internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53")}) + t.Cleanup(func() { internet.SkipDNSServers(nil) }) + preferGo, dial := net.DefaultResolver.PreferGo, net.DefaultResolver.Dial + saved := resolveOnOwn() + t.Cleanup(saved.restore) + if !net.DefaultResolver.PreferGo || net.DefaultResolver.Dial == nil { + t.Fatal("net.DefaultResolver is unchanged") + } + if _, err := net.DefaultResolver.Dial(context.Background(), "udp", "203.0.113.53:53"); err == nil { + t.Error("the TUN's DNS server was not skipped") + } + conn, err := net.DefaultResolver.Dial(context.Background(), "udp", "127.0.0.1:53") + if err != nil { + t.Fatal(err) + } + conn.Close() + saved.restore() + if net.DefaultResolver.PreferGo != preferGo || (net.DefaultResolver.Dial == nil) != (dial == nil) { + t.Error("net.DefaultResolver is not restored") + } +} + +// TestTunOnlyDNS checks that a DNS server another interface uses as well is +// not skipped, while one of the TUN alone is. +func TestTunOnlyDNS(t *testing.T) { + adapters, err := winipcfg.GetAdaptersAddresses(windows.AF_UNSPEC, winipcfg.GAAFlagIncludeGateways) + if err != nil { + t.Fatal(err) + } + var other netip.Addr + for _, adapter := range adapters { + if adapter.OperStatus == winipcfg.IfOperStatusUp && adapter.FirstGatewayAddress != nil && adapter.FirstDNSServerAddress != nil { + other, _ = netip.AddrFromSlice(adapter.FirstDNSServerAddress.Address.IP()) + other = other.Unmap() + break + } + } + if !other.IsValid() { + t.Skip("no interface with a gateway and a DNS server") + } + tunOnly := netip.MustParseAddr("203.0.113.53") + // LUID 0 is no interface, so every one counts as another. + got, err := tunOnlyDNS(0, []netip.Addr{other, tunOnly}) + if err != nil { + t.Fatal(err) + } + if !slices.Equal(got, []netip.Addr{tunOnly}) { + t.Errorf("got %v, want [%v]", got, tunOnly) + } +} + +func TestFlushDNSCache(t *testing.T) { + if err := flushDNSCache(); err != nil { + t.Fatal(err) + } +} diff --git a/proxy/tun/udp_fullcone.go b/proxy/tun/udp_fullcone.go index 1a1cb0843..480bb4bb7 100644 --- a/proxy/tun/udp_fullcone.go +++ b/proxy/tun/udp_fullcone.go @@ -78,11 +78,11 @@ func (u *udpConnectionHandler) HandlePacket(src net.Destination, dst net.Destina } } -func (u *udpConnectionHandler) connectionFinished(src net.Destination) { +func (u *udpConnectionHandler) connectionFinished(conn *udpConn) { u.Lock() - conn, found := u.udpConns[src] - if found { - delete(u.udpConns, src) + // Close runs twice per flow; a newer conn may already own this src. + if u.udpConns[conn.src] == conn { + delete(u.udpConns, conn.src) close(conn.egress) } u.Unlock() @@ -161,7 +161,7 @@ func (c *udpConn) Write(p []byte) (int, error) { } func (c *udpConn) Close() error { - c.handler.connectionFinished(c.src) + c.handler.connectionFinished(c) return nil } diff --git a/proxy/vless/account.go b/proxy/vless/account.go index 2f617149d..b773faa28 100644 --- a/proxy/vless/account.go +++ b/proxy/vless/account.go @@ -12,7 +12,7 @@ import ( func (a *Account) AsAccount() (protocol.Account, error) { id, err := uuid.ParseString(a.Id) if err != nil { - return nil, errors.New("failed to parse ID").Base(err).AtError() + return nil, errors.New("failed to parse ID").Base(err) } return &MemoryAccount{ ID: protocol.NewID(id), diff --git a/proxy/vless/inbound/inbound.go b/proxy/vless/inbound/inbound.go index 2d10a0060..d6a9d8efa 100644 --- a/proxy/vless/inbound/inbound.go +++ b/proxy/vless/inbound/inbound.go @@ -61,10 +61,10 @@ func init() { for _, user := range c.Users { u, err := user.ToMemoryUser() if err != nil { - return nil, errors.New("failed to get VLESS user").Base(err).AtError() + return nil, errors.New("failed to get VLESS user").Base(err) } if err := validator.Add(u); err != nil { - return nil, errors.New("failed to initiate user").Base(err).AtError() + return nil, errors.New("failed to initiate user").Base(err) } } @@ -110,7 +110,7 @@ func New(ctx context.Context, config *Config, dc dns.Client, validator vless.Val } handler.decryption = &encryption.ServerInstance{} if err := handler.decryption.Init(nfsSKeysBytes, config.XorMode, config.SecondsFrom, config.SecondsTo, config.Padding); err != nil { - return nil, errors.New("failed to use decryption").Base(err).AtError() + return nil, errors.New("failed to use decryption").Base(err) } } @@ -128,7 +128,7 @@ func New(ctx context.Context, config *Config, dc dns.Client, validator vless.Val /* if fb.Path != "" { if r, err := regexp.Compile(fb.Path); err != nil { - return nil, errors.New("invalid path regexp").Base(err).AtError() + return nil, errors.New("invalid path regexp").Base(err) } else { handler.regexps[fb.Path] = r } @@ -274,13 +274,13 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s if h.decryption != nil { var err error if connection, err = h.decryption.Handshake(connection, nil); err != nil { - return errors.New("ML-KEM-768 handshake failed").Base(err).AtInfo() + return errors.New("ML-KEM-768 handshake failed").Base(err) } } sessionPolicy := h.policyManager.ForLevel(0) if err := connection.SetReadDeadline(time.Now().Add(sessionPolicy.Timeouts.Handshake)); err != nil { - return errors.New("unable to set read deadline").Base(err).AtWarning() + return errors.New("unable to set read deadline").Base(err) } first := buf.FromBytes(make([]byte, buf.Size)) @@ -352,7 +352,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s } apfb := napfb[name] if apfb == nil { - return errors.New(`failed to find the default "name" config`).AtWarning() + return errors.New(`failed to find the default "name" config`) } if apfb[alpn] == nil { @@ -360,7 +360,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s } pfb := apfb[alpn] if pfb == nil { - return errors.New(`failed to find the default "alpn" config`).AtWarning() + return errors.New(`failed to find the default "alpn" config`) } path := "" @@ -369,7 +369,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s if lines := bytes.Split(firstBytes, []byte{'\r', '\n'}); len(lines) > 1 { if s := bytes.Split(lines[0], []byte{' '}); len(s) == 3 { if len(s[0]) < 8 && len(s[1]) > 0 && len(s[2]) == 8 { - errors.New("realPath = " + string(s[1])).AtInfo().WriteToLog(sid) + errors.New("realPath = " + string(s[1])).WriteToLog(sid) for _, fb := range pfb { if fb.Path != "" && h.regexps[fb.Path].Match(s[1]) { path = fb.Path @@ -409,7 +409,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s } fb := pfb[path] if fb == nil { - return errors.New(`failed to find the default "path" config`).AtWarning() + return errors.New(`failed to find the default "path" config`) } ctx, cancel := context.WithCancel(ctx) @@ -425,7 +425,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s } return nil }); err != nil { - return errors.New("failed to dial to " + fb.Dest).Base(err).AtWarning() + return errors.New("failed to dial to " + fb.Dest).Base(err) } defer conn.Close() @@ -485,11 +485,11 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s pro.Write([]byte{byte(p1 >> 8), byte(p1), byte(p2 >> 8), byte(p2)}) } if err := serverWriter.WriteMultiBuffer(buf.MultiBuffer{pro}); err != nil { - return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err).AtWarning() + return errors.New("failed to set PROXY protocol v", fb.Xver).Base(err) } } if err := buf.Copy(reader, serverWriter, buf.UpdateActivity(timer)); err != nil { - return errors.New("failed to fallback request payload").Base(err).AtInfo() + return errors.New("failed to fallback request payload").Base(err) } return nil } @@ -499,7 +499,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s getResponse := func() error { defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly) if err := buf.Copy(serverReader, writer, buf.UpdateActivity(timer)); err != nil { - return errors.New("failed to deliver response payload").Base(err).AtInfo() + return errors.New("failed to deliver response payload").Base(err) } return nil } @@ -507,7 +507,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s if err := task.Run(ctx, task.OnSuccess(postRequest, task.Close(serverWriter)), task.OnSuccess(getResponse, task.Close(writer))); err != nil { common.Interrupt(serverReader) common.Interrupt(serverWriter) - return errors.New("fallback ends").Base(err).AtInfo() + return errors.New("fallback ends").Base(err) } return nil } @@ -519,7 +519,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s Status: log.AccessRejected, Reason: err, }) - err = errors.New("invalid request from ", connection.RemoteAddr()).Base(err).AtInfo() + err = errors.New("invalid request from ", connection.RemoteAddr()).Base(err) } return err } @@ -555,7 +555,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s inbound.CanSpliceCopy = 2 switch request.Command { case protocol.RequestCommandUDP: - return errors.New(requestAddons.Flow + " doesn't support UDP").AtWarning() + return errors.New(requestAddons.Flow + " doesn't support UDP") case protocol.RequestCommandMux, protocol.RequestCommandRvs: inbound.CanSpliceCopy = 3 fallthrough // we will break Mux connections that contain TCP requests @@ -570,7 +570,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s p = uintptr(unsafe.Pointer(commonConn)) } else if tlsConn, ok := iConn.(*tls.Conn); ok { if tlsConn.ConnectionState().Version != gotls.VersionTLS13 { - return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, tlsConn.ConnectionState().Version).AtWarning() + return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, tlsConn.ConnectionState().Version) } t = reflect.TypeOf(tlsConn.Conn).Elem() p = uintptr(unsafe.Pointer(tlsConn.Conn)) @@ -578,7 +578,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s t = reflect.TypeOf(realityConn.Conn).Elem() p = uintptr(unsafe.Pointer(realityConn.Conn)) } else { - return errors.New("XTLS only supports TLS and REALITY directly for now.").AtWarning() + return errors.New("XTLS only supports TLS and REALITY directly for now.") } i, _ := t.FieldByName("input") r, _ := t.FieldByName("rawInput") @@ -586,15 +586,15 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s rawInput = (*bytes.Buffer)(unsafe.Pointer(p + r.Offset)) } } else { - return errors.New("account " + account.ID.String() + " is not able to use the flow " + requestAddons.Flow).AtWarning() + return errors.New("account " + account.ID.String() + " is not able to use the flow " + requestAddons.Flow) } case "": inbound.CanSpliceCopy = 3 if account.Flow == vless.XRV && (request.Command == protocol.RequestCommandTCP || isMuxAndNotXUDP(request, first)) { - return errors.New("account " + account.ID.String() + " is rejected since the client flow is empty. Note that the pure TLS proxy has certain TLS in TLS characters.").AtWarning() + return errors.New("account " + account.ID.String() + " is rejected since the client flow is empty. Note that the pure TLS proxy has certain TLS in TLS characters.") } default: - return errors.New("unknown request flow " + requestAddons.Flow).AtWarning() + return errors.New("unknown request flow " + requestAddons.Flow) } if request.Command != protocol.RequestCommandMux { @@ -617,7 +617,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s bufferWriter := buf.NewBufferedWriter(buf.NewWriter(connection)) if err := encoding.EncodeResponseHeader(bufferWriter, request, responseAddons); err != nil { - return errors.New("failed to encode response header").Base(err).AtWarning() + return errors.New("failed to encode response header").Base(err) } clientWriter := encoding.EncodeBodyAddons(bufferWriter, request, requestAddons, trafficState, false, ctx, connection, nil) bufferWriter.SetFlushNext() @@ -654,11 +654,11 @@ func (r *Reverse) Tag() string { func (r *Reverse) NewMux(ctx context.Context, link *transport.Link, observer features.Feature) error { muxClient, err := mux.NewClientWorker(*link, mux.ClientStrategy{}) 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 := reverse.NewPortalWorker(muxClient) if err != nil { - return errors.New("failed to create portal worker").Base(err).AtWarning() + return errors.New("failed to create portal worker").Base(err) } r.picker.AddWorker(worker) if burstObs, ok := observer.(extension.BurstObservatory); ok { diff --git a/proxy/vless/outbound/outbound.go b/proxy/vless/outbound/outbound.go index 6a013dae2..56a2fd04e 100644 --- a/proxy/vless/outbound/outbound.go +++ b/proxy/vless/outbound/outbound.go @@ -73,7 +73,7 @@ func New(ctx context.Context, config *Config) (*Handler, error) { } server, err := protocol.NewServerSpecFromPB(config.Vnext) if err != nil { - return nil, errors.New("failed to get server spec").Base(err).AtError() + return nil, errors.New("failed to get server spec").Base(err) } v := core.MustFromContext(ctx) @@ -93,7 +93,7 @@ func New(ctx context.Context, config *Config) (*Handler, error) { } handler.encryption = &encryption.ClientInstance{} if err := handler.encryption.Init(nfsPKeysBytes, a.XorMode, a.Seconds, a.Padding); err != nil { - return nil, errors.New("failed to use encryption").Base(err).AtError() + return nil, errors.New("failed to use encryption").Base(err) } } @@ -106,7 +106,7 @@ func New(ctx context.Context, config *Config) (*Handler, error) { if sc := a.Reverse.Sniffing; sc != nil && sc.Enabled { request, err := proxymanConfig.BuildSniffingRequest(sc) if err != nil { - return nil, errors.New("failed to build reverse sniffing request").Base(err).AtError() + return nil, errors.New("failed to build reverse sniffing request").Base(err) } rvsCtx = session.ContextWithContent(rvsCtx, &session.Content{ SniffingRequest: request, @@ -149,7 +149,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte outbounds := session.OutboundsFromContext(ctx) ob := outbounds[len(outbounds)-1] if !ob.Target.IsValid() && ob.Target.Address.String() != "v1.rvs.cool" { - return errors.New("target not specified").AtError() + return errors.New("target not specified") } ob.Name = "vless" @@ -178,7 +178,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte for { connTime := <-h.preConns if connTime == nil { - return errors.New("closed handler").AtWarning() + return errors.New("closed handler") } if time.Now().Before(connTime.Expire) { conn = connTime.Conn @@ -197,7 +197,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte } return nil }); err != nil { - return errors.New("failed to find an available destination").Base(err).AtWarning() + return errors.New("failed to find an available destination").Base(err) } } defer conn.Close() @@ -209,7 +209,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte if h.encryption != nil { var err error if conn, err = h.encryption.Handshake(conn); err != nil { - return errors.New("ML-KEM-768 handshake failed").Base(err).AtInfo() + return errors.New("ML-KEM-768 handshake failed").Base(err) } } @@ -223,7 +223,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte command = protocol.RequestCommandMux case "v1.rvs.cool": if target.Network != net.Network_Unknown { - return errors.New("nice try baby").AtError() + return errors.New("nice try baby") } command = protocol.RequestCommandRvs } @@ -256,7 +256,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte switch request.Command { case protocol.RequestCommandUDP: if !allowUDP443 && request.Port == 443 { - return errors.New("XTLS rejected UDP/443 traffic").AtInfo() + return errors.New("XTLS rejected UDP/443 traffic") } case protocol.RequestCommandMux: fallthrough // let server break Mux connections that contain TCP requests @@ -279,7 +279,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte t = reflect.TypeOf(realityConn.Conn).Elem() p = uintptr(unsafe.Pointer(realityConn.Conn)) } else { - return errors.New("XTLS only supports TLS and REALITY directly for now.").AtWarning() + return errors.New("XTLS only supports TLS and REALITY directly for now.") } i, _ := t.FieldByName("input") r, _ := t.FieldByName("rawInput") @@ -321,7 +321,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte bufferWriter := buf.NewBufferedWriter(buf.NewWriter(conn)) if err := encoding.EncodeRequestHeader(bufferWriter, request, requestAddons); err != nil { - return errors.New("failed to encode request header").Base(err).AtWarning() + return errors.New("failed to encode request header").Base(err) } // default: serverWriter := bufferWriter @@ -350,23 +350,23 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte } // Flush; bufferWriter.WriteMultiBuffer now is bufferWriter.writer.WriteMultiBuffer if err := bufferWriter.SetBuffered(false); err != nil { - return errors.New("failed to write A request payload").Base(err).AtWarning() + return errors.New("failed to write A request payload").Base(err) } if requestAddons.Flow == vless.XRV { if tlsConn, ok := iConn.(*tls.Conn); ok { if tlsConn.ConnectionState().Version != gotls.VersionTLS13 { - return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, tlsConn.ConnectionState().Version).AtWarning() + return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, tlsConn.ConnectionState().Version) } } else if utlsConn, ok := iConn.(*tls.UConn); ok { if utlsConn.ConnectionState().Version != utls.VersionTLS13 { - return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, utlsConn.ConnectionState().Version).AtWarning() + return errors.New(`failed to use `+requestAddons.Flow+`, found outer tls version `, utlsConn.ConnectionState().Version) } } } err := buf.Copy(clientReader, serverWriter, buf.UpdateActivity(timer)) if err != nil { - return errors.New("failed to transfer request payload").Base(err).AtInfo() + return errors.New("failed to transfer request payload").Base(err) } // Indicates the end of request payload. @@ -381,7 +381,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte responseAddons, err := encoding.DecodeResponseHeader(conn, request) if err != nil { - return errors.New("failed to decode response header").Base(err).AtInfo() + return errors.New("failed to decode response header").Base(err) } // default: serverReader := buf.NewReader(conn) @@ -405,7 +405,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte } if err != nil { - return errors.New("failed to transfer response payload").Base(err).AtInfo() + return errors.New("failed to transfer response payload").Base(err) } return nil @@ -416,7 +416,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte } if err := task.Run(ctx, postRequest, task.OnSuccess(getResponse, task.Close(clientWriter))); err != nil { - return errors.New("connection ends").Base(err).AtInfo() + return errors.New("connection ends").Base(err) } return nil diff --git a/proxy/vmess/account.go b/proxy/vmess/account.go index a7a1b78ef..a2b7210b1 100644 --- a/proxy/vmess/account.go +++ b/proxy/vmess/account.go @@ -49,7 +49,7 @@ func (a *MemoryAccount) ToProto() proto.Message { func (a *Account) AsAccount() (protocol.Account, error) { id, err := uuid.ParseString(a.Id) if err != nil { - return nil, errors.New("failed to parse ID").Base(err).AtError() + return nil, errors.New("failed to parse ID").Base(err) } protoID := protocol.NewID(id) var AuthenticatedLength, NoTerminationSignal bool diff --git a/proxy/vmess/encoding/client.go b/proxy/vmess/encoding/client.go index e5b38d391..2890fc7b1 100644 --- a/proxy/vmess/encoding/client.go +++ b/proxy/vmess/encoding/client.go @@ -209,7 +209,7 @@ func (c *ClientSession) DecodeResponseHeader(reader io.Reader) (*protocol.Respon defer buffer.Release() if _, err := buffer.ReadFullFrom(c.responseReader, 4); err != nil { - return nil, errors.New("failed to read response header").Base(err).AtWarning() + return nil, errors.New("failed to read response header").Base(err) } if buffer.Byte(0) != c.responseHeader { diff --git a/proxy/vmess/inbound/inbound.go b/proxy/vmess/inbound/inbound.go index e084a881e..03006ead1 100644 --- a/proxy/vmess/inbound/inbound.go +++ b/proxy/vmess/inbound/inbound.go @@ -227,7 +227,7 @@ func transferResponse(timer signal.ActivityUpdater, session *encoding.ServerSess func (h *Handler) Process(ctx context.Context, network net.Network, connection stat.Connection, dispatcher routing.Dispatcher) error { sessionPolicy := h.policyManager.ForLevel(0) if err := connection.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) } iConn := stat.TryUnwrapStatsConn(connection) @@ -247,7 +247,7 @@ func (h *Handler) Process(ctx context.Context, network net.Network, connection s Status: log.AccessRejected, Reason: err, }) - err = errors.New("invalid request from ", connection.RemoteAddr()).Base(err).AtInfo() + err = errors.New("invalid request from ", connection.RemoteAddr()).Base(err) } return err } diff --git a/proxy/vmess/outbound/outbound.go b/proxy/vmess/outbound/outbound.go index db5f94abb..a475dac92 100644 --- a/proxy/vmess/outbound/outbound.go +++ b/proxy/vmess/outbound/outbound.go @@ -60,7 +60,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte outbounds := session.OutboundsFromContext(ctx) ob := outbounds[len(outbounds)-1] if !ob.Target.IsValid() { - return errors.New("target not specified").AtError() + return errors.New("target not specified") } ob.Name = "vmess" ob.CanSpliceCopy = 3 @@ -78,7 +78,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte return nil }) if err != nil { - return errors.New("failed to find an available destination").Base(err).AtWarning() + return errors.New("failed to find an available destination").Base(err) } defer conn.Close() @@ -154,7 +154,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte writer := buf.NewBufferedWriter(buf.NewWriter(conn)) if err := session.EncodeRequestHeader(request, writer); err != nil { - return errors.New("failed to encode request").Base(err).AtWarning() + return errors.New("failed to encode request").Base(err) } bodyWriter, err := session.EncodeRequestBody(request, writer) diff --git a/proxy/wireguard/bind.go b/proxy/wireguard/bind.go index 0e71e1539..fda540d4d 100644 --- a/proxy/wireguard/bind.go +++ b/proxy/wireguard/bind.go @@ -52,9 +52,12 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err case <-ch: default: errors.LogErrorInner(context.Background(), err, "unexpected closed") - if b.downFunc != nil { + b.mu.Lock() + downFunc := b.downFunc + b.mu.Unlock() + if downFunc != nil { go func() { - common.Must(b.downFunc()) + common.Must(downFunc()) }() } } @@ -76,6 +79,13 @@ func (b *bind) Open(port uint16) (fns []conn.ReceiveFunc, actualPort uint16, err }, uint16(c.LocalAddr().(*net.UDPAddr).Port), nil } +// setDownFunc sets downFunc after the device is created, since the device may already be using the bind. +func (b *bind) setDownFunc(f func() error) { + b.mu.Lock() + defer b.mu.Unlock() + b.downFunc = f +} + func (b *bind) Close() error { b.mu.Lock() defer b.mu.Unlock() diff --git a/proxy/wireguard/client.go b/proxy/wireguard/client.go index 9cfff7265..c47e3c8b1 100644 --- a/proxy/wireguard/client.go +++ b/proxy/wireguard/client.go @@ -27,7 +27,6 @@ import ( "github.com/xtls/xray-core/features/stats" "github.com/xtls/xray-core/transport" "github.com/xtls/xray-core/transport/internet" - "github.com/xtls/xray-core/transport/internet/finalmask" "golang.zx2c4.com/wireguard/device" ) @@ -200,7 +199,7 @@ func (h *Handler) Process(ctx context.Context, link *transport.Link, dialer inte } defer conn.Close() c := &UDPConnClient{ - PacketConn: conn.(*internet.PacketConnWrapper).PacketConn, + PacketConn: conn.(*net.PacketConnWrapper).PacketConn, Dest: conn.RemoteAddr().(*net.UDPAddr), } reader = c @@ -264,14 +263,14 @@ func (h *Handler) init(ctx context.Context) error { if err != nil { return nil, errors.New("failed to dial to dest").Base(err) } - pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn + pktConn = conn.(*net.PacketConnWrapper).PacketConn } else { conn, err := internet.DialSystem(ctx, dest, h.streamSettings.SocketSettings) if err != nil { return nil, errors.New("failed to dial to dest").Base(err) } switch c := conn.(type) { - case *internet.PacketConnWrapper: + case *net.PacketConnWrapper: pktConn = c.PacketConn case *cnc.Connection: pktConn = &internet.FakePacketConn{Conn: c} @@ -288,7 +287,13 @@ func (h *Handler) init(ctx context.Context) error { } return pktConn, nil } - bind := &bind{} + // device.NewDevice may use the bind right away (Up -> BindUpdate -> Open), + // so everything it reads must be set before creating the device. + bind := &bind{ + resolveFunc: resolveFunc, + listenFunc: listenFunc, + reserved: h.conf.Reserved, + } logger := &device.Logger{ Verbosef: func(format string, args ...any) { log.Record(&log.GeneralMessage{ @@ -304,10 +309,7 @@ func (h *Handler) init(ctx context.Context) error { }, } dev := device.NewDevice(h.tun, bind, logger) - bind.resolveFunc = resolveFunc - bind.listenFunc = listenFunc - bind.downFunc = dev.Down - bind.reserved = h.conf.Reserved + bind.setDownFunc(dev.Down) var cfg strings.Builder cfg.WriteString("private_key=" + h.conf.SecretKey + "\n") for _, peer := range h.conf.Peers { diff --git a/proxy/wireguard/netstack.go b/proxy/wireguard/netstack.go index b54961c56..cc04f1580 100644 --- a/proxy/wireguard/netstack.go +++ b/proxy/wireguard/netstack.go @@ -21,7 +21,7 @@ import ( "syscall" "time" - "github.com/xtls/xray-core/transport/internet" + xnet "github.com/xtls/xray-core/common/net" "golang.zx2c4.com/wireguard/tun" "golang.org/x/net/dns/dnsmessage" @@ -220,7 +220,7 @@ func (tun *netTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, error if err != nil { return nil, err } - return &internet.PacketConnWrapper{ + return &xnet.PacketConnWrapper{ PacketConn: conn, Dest: net.UDPAddrFromAddrPort(raddr), }, nil diff --git a/proxy/wireguard/server.go b/proxy/wireguard/server.go index a9a0a73ca..5eef90e7a 100644 --- a/proxy/wireguard/server.go +++ b/proxy/wireguard/server.go @@ -113,7 +113,7 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) { users.Store(user.Account.(*MemoryAccount).Pub, user) } - return &Server{ + s := &Server{ conf: conf, ctx: core.ToBackgroundDetachedContext(ctx), policyManager: p, @@ -131,7 +131,10 @@ func NewServer(ctx context.Context, conf *DeviceConfig) (*Server, error) { pub: pub, users: users, - }, nil + } + // Install the stack's protocol handlers before the device can deliver packets to it (Start -> dev.Up). + CreateForwarder(stack, s.HandleConnection) + return s, nil } func (s *Server) AddUser(ctx context.Context, user *protocol.MemoryUser) error { @@ -320,7 +323,6 @@ func (s *Server) Start() error { return err } s.dev = dev - createForwarder(s.stack, s.HandleConnection) return nil } diff --git a/proxy/wireguard/tun.go b/proxy/wireguard/tun.go index 68bad3ac6..b4bcdc371 100644 --- a/proxy/wireguard/tun.go +++ b/proxy/wireguard/tun.go @@ -49,7 +49,7 @@ func CalculateInterfaceName(name string) (tunName string) { return } -func createForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.Destination)) { +func CreateForwarder(gstack *stack.Stack, handler func(conn net.Conn, dest net.Destination)) { gstack.SetPromiscuousMode(1, true) gstack.SetSpoofing(1, true) diff --git a/proxy/wireguard/tun_linux.go b/proxy/wireguard/tun_linux.go index eb0175476..08a7e4dd5 100644 --- a/proxy/wireguard/tun_linux.go +++ b/proxy/wireguard/tun_linux.go @@ -16,7 +16,7 @@ import ( "github.com/vishvananda/netlink" "github.com/xtls/xray-core/common/errors" - "github.com/xtls/xray-core/transport/internet" + xnet "github.com/xtls/xray-core/common/net" "golang.zx2c4.com/wireguard/tun" ) @@ -263,7 +263,7 @@ func (tun *kernelTun) DialUDPAddrPort(laddr, raddr netip.AddrPort) (net.Conn, er if err != nil { return nil, err } - return &internet.PacketConnWrapper{ + return &xnet.PacketConnWrapper{ PacketConn: conn, Dest: net.UDPAddrFromAddrPort(raddr), }, nil diff --git a/testing/scenarios/masque_test.go b/testing/scenarios/masque_test.go index 1b203a88a..95526a222 100644 --- a/testing/scenarios/masque_test.go +++ b/testing/scenarios/masque_test.go @@ -1,19 +1,29 @@ package scenarios import ( + "bufio" + "bytes" "context" + "crypto/rand" gotls "crypto/tls" "crypto/x509" + "encoding/binary" go_errors "errors" "io" "net/http" "net/netip" + "net/url" + "strconv" + "strings" + "sync" "sync/atomic" "testing" "time" "github.com/apernet/quic-go" "github.com/apernet/quic-go/http3" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" "golang.org/x/sync/errgroup" "gvisor.dev/gvisor/pkg/tcpip" "gvisor.dev/gvisor/pkg/tcpip/adapters/gonet" @@ -22,6 +32,7 @@ import ( "github.com/xtls/xray-core/app/log" "github.com/xtls/xray-core/app/proxyman" + "github.com/xtls/xray-core/app/router" "github.com/xtls/xray-core/common" clog "github.com/xtls/xray-core/common/log" "github.com/xtls/xray-core/common/net" @@ -30,6 +41,7 @@ import ( "github.com/xtls/xray-core/common/serial" core "github.com/xtls/xray-core/core" "github.com/xtls/xray-core/proxy/dokodemo" + "github.com/xtls/xray-core/proxy/freedom" "github.com/xtls/xray-core/proxy/masque" "github.com/xtls/xray-core/proxy/wireguard" "github.com/xtls/xray-core/testing/servers/tcp" @@ -49,10 +61,10 @@ var ( const ( masqueEchoPort = 7 - masqueAuthorization = "Basic dTpw" + masqueAuthorization = "Basic dUBleGFtcGxlLmNvbTpw" ) -func startMasqueServer(t *testing.T) (net.Port, [32]byte) { +func startMasqueServer(t *testing.T, h2 bool) (net.Port, [32]byte) { dev, _, gstack, err := wireguard.CreateNetTUN([]netip.Addr{masqueServerV4, masqueServerV6}, nil, transmasque.MinPacketSize, false) common.Must(err) t.Cleanup(func() { dev.Close() }) @@ -180,6 +192,13 @@ func startMasqueServer(t *testing.T) (net.Port, [32]byte) { Certificates: []gotls.Certificate{{Certificate: [][]byte{certificate.Certificate}, PrivateKey: key}}, NextProtos: []string{http3.NextProtoH3}, } + if h2 { + tlsConfig.NextProtos = []string{http2.NextProtoTLS} + ln := common.Must2(gotls.Listen("tcp", "127.0.0.1:0", tlsConfig)) + t.Cleanup(func() { ln.Close() }) + go serveHTTP2(ln, http.HandlerFunc(handler)) + return net.Port(ln.Addr().(*net.TCPAddr).Port), certHash + } pktConn := common.Must2(net.ListenUDP("udp", &net.UDPAddr{IP: net.LocalHostIP.IP()})) tr := &quic.Transport{Conn: pktConn} ln := common.Must2(tr.ListenEarly(tlsConfig, &quic.Config{EnableDatagrams: true, InitialPacketSize: 1350})) @@ -195,26 +214,239 @@ func startMasqueServer(t *testing.T) (net.Port, [32]byte) { return net.Port(pktConn.LocalAddr().(*net.UDPAddr).Port), certHash } -func TestMasque(t *testing.T) { - serverPort, certHash := startMasqueServer(t) +func serveHTTP2(ln net.Listener, handler http.Handler) { + for { + conn, err := ln.Accept() + if err != nil { + return + } + go serveHTTP2Conn(conn, handler) + } +} - tcpPort := tcp.PickPort() - tcp6Port := tcp.PickPort() - udpPort := udp.PickPort() - dokodemoTo := func(port net.Port, addr netip.Addr, network net.Network) *core.InboundHandlerConfig { - return &core.InboundHandlerConfig{ - ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{ - PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(port)}}, - Listen: net.NewIPOrDomain(net.LocalHostIP), - }), - ProxySettings: serial.ToTypedMessage(&dokodemo.Config{ - RewriteAddress: net.NewIPOrDomain(net.IPAddress(addr.AsSlice())), - RewritePort: masqueEchoPort, - AllowedNetworks: []net.Network{network}, - }), +type http2ServerConn struct { + mu sync.Mutex + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder +} + +func (c *http2ServerConn) write(f func(*http2.Framer) error) error { + c.mu.Lock() + defer c.mu.Unlock() + return f(c.fr) +} + +func (c *http2ServerConn) writeHeaders(streamID uint32, status int, header http.Header) error { + c.mu.Lock() + defer c.mu.Unlock() + c.hbuf.Reset() + c.henc.WriteField(hpack.HeaderField{Name: ":status", Value: strconv.Itoa(status)}) + for k, vv := range header { + for _, v := range vv { + c.henc.WriteField(hpack.HeaderField{Name: strings.ToLower(k), Value: v}) } } - clientConfig := &core.Config{ + return c.fr.WriteHeaders(http2.HeadersFrameParam{StreamID: streamID, BlockFragment: c.hbuf.Bytes(), EndHeaders: true}) +} + +func (c *http2ServerConn) writeData(streamID uint32, endStream bool, data []byte) error { + c.mu.Lock() + defer c.mu.Unlock() + for { + n := min(len(data), 16384) + if err := c.fr.WriteData(streamID, endStream && n == len(data), data[:n]); err != nil { + return err + } + if data = data[n:]; len(data) == 0 { + return nil + } + } +} + +func serveHTTP2Conn(conn net.Conn, handler http.Handler) { + defer conn.Close() + br := bufio.NewReader(conn) + preface := make([]byte, len(http2.ClientPreface)) + if _, err := io.ReadFull(br, preface); err != nil || string(preface) != http2.ClientPreface { + return + } + sc := &http2ServerConn{fr: http2.NewFramer(conn, br)} + sc.henc = hpack.NewEncoder(&sc.hbuf) + sc.fr.ReadMetaHeaders = hpack.NewDecoder(4096, nil) + if err := sc.write(func(fr *http2.Framer) error { + if err := fr.WriteSettings( + http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: 1 << 30}, + ); err != nil { + return err + } + return fr.WriteWindowUpdate(0, 1<<30) + }); err != nil { + return + } + + bodies := make(map[uint32]*io.PipeWriter) + defer func() { + for _, body := range bodies { + body.Close() + } + }() + for { + f, err := sc.fr.ReadFrame() + if err != nil { + return + } + switch f := f.(type) { + case *http2.SettingsFrame: + if !f.IsAck() { + err = sc.write((*http2.Framer).WriteSettingsAck) + } + case *http2.PingFrame: + if !f.IsAck() { + err = sc.write(func(fr *http2.Framer) error { return fr.WritePing(true, f.Data) }) + } + case *http2.MetaHeadersFrame: + u, err := url.ParseRequestURI(f.PseudoValue("path")) + if err != nil { + return + } + pr, pw := io.Pipe() + bodies[f.StreamID] = pw + req := &http.Request{ + Method: f.PseudoValue("method"), + URL: u, + Proto: "HTTP/2.0", + ProtoMajor: 2, + Header: http.Header{}, + Host: f.PseudoValue("authority"), + Body: pr, + } + for _, hf := range f.RegularFields() { + req.Header.Add(hf.Name, hf.Value) + } + if protocol := f.PseudoValue("protocol"); protocol != "" { + req.Header.Set(":protocol", protocol) + } + streamID := f.StreamID + w := &http2ResponseWriter{conn: sc, streamID: streamID, header: http.Header{}} + go func() { + handler.ServeHTTP(w, req) + w.WriteHeader(http.StatusOK) + sc.writeData(streamID, true, nil) + }() + case *http2.DataFrame: + if body := bodies[f.StreamID]; body != nil { + if _, err := body.Write(f.Data()); err != nil || f.StreamEnded() { + body.Close() + delete(bodies, f.StreamID) + } + } + case *http2.RSTStreamFrame: + if body := bodies[f.StreamID]; body != nil { + body.CloseWithError(http2.StreamError{StreamID: f.StreamID, Code: f.ErrCode}) + delete(bodies, f.StreamID) + } + } + if err != nil { + return + } + } +} + +type http2ResponseWriter struct { + conn *http2ServerConn + streamID uint32 + header http.Header + wroteHeader bool +} + +func (w *http2ResponseWriter) Header() http.Header { return w.header } + +func (w *http2ResponseWriter) WriteHeader(code int) { + if !w.wroteHeader { + w.wroteHeader = true + w.conn.writeHeaders(w.streamID, code, w.header) + } +} + +func (w *http2ResponseWriter) Write(b []byte) (int, error) { + w.WriteHeader(http.StatusOK) + if err := w.conn.writeData(w.streamID, false, b); err != nil { + return 0, err + } + return len(b), nil +} + +func (w *http2ResponseWriter) Flush() {} + +func TestMasque(t *testing.T) { + testMasque(t, false) +} + +func TestMasqueHTTP2(t *testing.T) { + testMasque(t, true) +} + +func masqueDokodemo(port net.Port, addr netip.Addr, network net.Network) *core.InboundHandlerConfig { + return &core.InboundHandlerConfig{ + ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{ + PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(port)}}, + Listen: net.NewIPOrDomain(net.LocalHostIP), + }), + ProxySettings: serial.ToTypedMessage(&dokodemo.Config{ + RewriteAddress: net.NewIPOrDomain(net.IPAddress(addr.AsSlice())), + RewritePort: masqueEchoPort, + AllowedNetworks: []net.Network{network}, + }), + } +} + +func masqueStreamSettings(tlsConfig *tls.Config, config *transmasque.Config) *internet.StreamConfig { + return &internet.StreamConfig{ + ProtocolName: "masque", + TransportSettings: []*internet.TransportConfig{ + { + ProtocolName: "masque", + Settings: serial.ToTypedMessage(config), + }, + }, + SecurityType: serial.GetMessageType(&tls.Config{}), + SecuritySettings: []*serial.TypedMessage{serial.ToTypedMessage(tlsConfig)}, + } +} + +func masqueClientTLS(certHash [32]byte, alpn ...string) *tls.Config { + return &tls.Config{ + ServerName: "localhost", + PinnedPeerCertSha256: [][]byte{certHash[:]}, + NextProtocol: alpn, + } +} + +func masqueOutbound(serverPort net.Port, certHash [32]byte, h2 bool, authorization string) *core.OutboundHandlerConfig { + tlsConfig := masqueClientTLS(certHash) + if h2 { + tlsConfig.NextProtocol = []string{http2.NextProtoTLS} + } + return &core.OutboundHandlerConfig{ + ProxySettings: serial.ToTypedMessage(&masque.ClientConfig{ + Server: &protocol.ServerEndpoint{ + Address: net.NewIPOrDomain(net.LocalHostIP), + Port: uint32(serverPort), + }, + }), + SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{ + StreamSettings: masqueStreamSettings(tlsConfig, &transmasque.Config{ + Path: transmasque.DefaultPath, + Headers: map[string]string{"Authorization": authorization}, + }), + }), + } +} + +func masqueClientConfig(serverPort net.Port, certHash [32]byte, h2 bool, authorization string, tcpPort, tcp6Port, udpPort net.Port, v4, v6 netip.Addr) *core.Config { + return &core.Config{ App: []*serial.TypedMessage{ serial.ToTypedMessage(&log.Config{ ErrorLogLevel: clog.Severity_Debug, @@ -222,47 +454,17 @@ func TestMasque(t *testing.T) { }), }, Inbound: []*core.InboundHandlerConfig{ - dokodemoTo(tcpPort, masqueServerV4, net.Network_TCP), - dokodemoTo(tcp6Port, masqueServerV6, net.Network_TCP), - dokodemoTo(udpPort, masqueServerV4, net.Network_UDP), + masqueDokodemo(tcpPort, v4, net.Network_TCP), + masqueDokodemo(tcp6Port, v6, net.Network_TCP), + masqueDokodemo(udpPort, v4, net.Network_UDP), }, Outbound: []*core.OutboundHandlerConfig{ - { - ProxySettings: serial.ToTypedMessage(&masque.ClientConfig{ - Server: &protocol.ServerEndpoint{ - Address: net.NewIPOrDomain(net.LocalHostIP), - Port: uint32(serverPort), - }, - }), - SenderSettings: serial.ToTypedMessage(&proxyman.SenderConfig{ - StreamSettings: &internet.StreamConfig{ - ProtocolName: "masque", - TransportSettings: []*internet.TransportConfig{ - { - ProtocolName: "masque", - Settings: serial.ToTypedMessage(&transmasque.Config{ - Path: transmasque.DefaultPath, - Headers: map[string]string{"Authorization": masqueAuthorization}, - }), - }, - }, - SecurityType: serial.GetMessageType(&tls.Config{}), - SecuritySettings: []*serial.TypedMessage{ - serial.ToTypedMessage(&tls.Config{ - ServerName: "localhost", - PinnedPeerCertSha256: [][]byte{certHash[:]}, - }), - }, - }, - }), - }, + masqueOutbound(serverPort, certHash, h2, authorization), }, } +} - servers, err := InitializeServerConfigs(clientConfig) - common.Must(err) - defer CloseAllServers(servers) - +func testMasqueTraffic(t *testing.T, tcpPort, tcp6Port, udpPort net.Port) { var errg errgroup.Group for range 3 { errg.Go(testTCPConn(tcpPort, 1024*1024, time.Second*20)) @@ -273,3 +475,215 @@ func TestMasque(t *testing.T) { t.Error(err) } } + +func testMasque(t *testing.T, h2 bool) { + serverPort, certHash := startMasqueServer(t, h2) + + tcpPort := tcp.PickPort() + tcp6Port := tcp.PickPort() + udpPort := udp.PickPort() + clientConfig := masqueClientConfig(serverPort, certHash, h2, masqueAuthorization, tcpPort, tcp6Port, udpPort, masqueServerV4, masqueServerV6) + + servers, err := InitializeServerConfigs(clientConfig) + common.Must(err) + defer CloseAllServers(servers) + + testMasqueTraffic(t, tcpPort, tcp6Port, udpPort) +} + +func masqueServerInbound(serverPort net.Port, certificate *tls.Certificate, alpn ...string) *core.InboundHandlerConfig { + return &core.InboundHandlerConfig{ + ReceiverSettings: serial.ToTypedMessage(&proxyman.ReceiverConfig{ + PortList: &net.PortList{Range: []*net.PortRange{net.SinglePortRange(serverPort)}}, + Listen: net.NewIPOrDomain(net.LocalHostIP), + StreamSettings: masqueStreamSettings(&tls.Config{ + Certificate: []*tls.Certificate{certificate}, + NextProtocol: alpn, + }, &transmasque.Config{Path: transmasque.DefaultPath}), + }), + ProxySettings: serial.ToTypedMessage(&masque.ServerConfig{ + Users: []*protocol.User{{ + Email: "u@example.com", + Account: serial.ToTypedMessage(&masque.Account{Password: "p"}), + }}, + Address: []string{"10.14.0.1/24", "fd14::1/64"}, + }), + } +} + +func masqueServerConfig(serverPort net.Port, certificate *tls.Certificate, h2 bool, tcpDest, udpDest net.Destination) *core.Config { + var alpn []string + if h2 { + alpn = []string{http2.NextProtoTLS} + } + redirect := func(tag string, dest net.Destination) *core.OutboundHandlerConfig { + return &core.OutboundHandlerConfig{ + Tag: tag, + ProxySettings: serial.ToTypedMessage(&freedom.Config{ + DestinationOverride: &freedom.DestinationOverride{ + Server: &protocol.ServerEndpoint{ + Address: net.NewIPOrDomain(dest.Address), + Port: uint32(dest.Port), + }, + }, + FinalRules: []*freedom.FinalRuleConfig{{Action: freedom.RuleAction_Allow}}, + }), + } + } + return &core.Config{ + App: []*serial.TypedMessage{ + serial.ToTypedMessage(&log.Config{ + ErrorLogLevel: clog.Severity_Debug, + ErrorLogType: log.LogType_Console, + }), + serial.ToTypedMessage(&router.Config{ + Rule: []*router.RoutingRule{ + {Networks: []net.Network{net.Network_TCP}, TargetTag: &router.RoutingRule_Tag{Tag: "tcp"}}, + {Networks: []net.Network{net.Network_UDP}, TargetTag: &router.RoutingRule_Tag{Tag: "udp"}}, + }, + }), + }, + Inbound: []*core.InboundHandlerConfig{ + masqueServerInbound(serverPort, certificate, alpn...), + }, + Outbound: []*core.OutboundHandlerConfig{ + redirect("tcp", tcpDest), + redirect("udp", udpDest), + }, + } +} + +func testMasqueServer(t *testing.T, h2 bool, authorization string) error { + tcpServer := tcp.Server{MsgProcessor: xor} + tcpDest, err := tcpServer.Start() + common.Must(err) + defer tcpServer.Close() + udpServer := udp.Server{MsgProcessor: xor} + udpDest, err := udpServer.Start() + common.Must(err) + defer udpServer.Close() + + ct, ctHash := cert.MustGenerate(nil, cert.CommonName("localhost")) + serverPort := udp.PickPort() + if h2 { + serverPort = tcp.PickPort() + } + tcpPort := tcp.PickPort() + tcp6Port := tcp.PickPort() + udpPort := udp.PickPort() + servers, err := InitializeServerConfigs( + masqueServerConfig(serverPort, tls.ParseCertificate(ct), h2, tcpDest, udpDest), + masqueClientConfig(serverPort, ctHash, h2, authorization, tcpPort, tcp6Port, udpPort, netip.MustParseAddr("192.0.2.1"), netip.MustParseAddr("2001:db8::1")), + ) + common.Must(err) + defer CloseAllServers(servers) + + if authorization != masqueAuthorization { + return testTCPConn(tcpPort, 1024, time.Second*5)() + } + testMasqueTraffic(t, tcpPort, tcp6Port, udpPort) + return nil +} + +func TestMasqueServer(t *testing.T) { + testMasqueServer(t, false, masqueAuthorization) +} + +func TestMasqueServerHTTP2(t *testing.T) { + testMasqueServer(t, true, masqueAuthorization) +} + +func TestMasqueServerRejectsWrongPassword(t *testing.T) { + for _, h2 := range []bool{false, true} { + if err := testMasqueServer(t, h2, "Basic dUBleGFtcGxlLmNvbTp3cm9uZw=="); err == nil { + t.Errorf("a wrong password got through (h2: %v)", h2) + } + } +} + +func masqueIPPacket(src, dst netip.Addr, payload []byte) []byte { + if src.Is4() { + p := make([]byte, 20, 20+len(payload)) + p[0] = 0x45 + binary.BigEndian.PutUint16(p[2:], uint16(20+len(payload))) + p[8] = 64 + p[9] = 253 + copy(p[12:], src.AsSlice()) + copy(p[16:], dst.AsSlice()) + return append(p, payload...) + } + p := make([]byte, 40, 40+len(payload)) + p[0] = 0x60 + binary.BigEndian.PutUint16(p[4:], uint16(len(payload))) + p[6] = 253 + p[7] = 64 + copy(p[8:], src.AsSlice()) + copy(p[24:], dst.AsSlice()) + return append(p, payload...) +} + +func masqueIPAddrs(p []byte) (src, dst netip.Addr) { + if p[0]>>4 == 4 { + return netip.AddrFrom4([4]byte(p[12:16])), netip.AddrFrom4([4]byte(p[16:20])) + } + return netip.AddrFrom16([16]byte(p[8:24])), netip.AddrFrom16([16]byte(p[24:40])) +} + +func TestMasqueServerClientToClient(t *testing.T) { + ct, ctHash := cert.MustGenerate(nil, cert.CommonName("localhost")) + serverPort := udp.PickPort() + servers, err := InitializeServerConfigs(&core.Config{ + Inbound: []*core.InboundHandlerConfig{ + masqueServerInbound(serverPort, tls.ParseCertificate(ct), http3.NextProtoH3, http2.NextProtoTLS), + }, + Outbound: []*core.OutboundHandlerConfig{ + {ProxySettings: serial.ToTypedMessage(&freedom.Config{})}, + }, + }) + common.Must(err) + defer CloseAllServers(servers) + + dial := func(alpn ...string) *transmasque.Conn { + streamSettings, err := internet.ToMemoryStreamConfig(masqueStreamSettings(masqueClientTLS(ctHash, alpn...), &transmasque.Config{ + Path: transmasque.DefaultPath, + Headers: map[string]string{"Authorization": masqueAuthorization}, + })) + common.Must(err) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + conn, err := transmasque.Dial(ctx, net.TCPDestination(net.LocalHostIP, serverPort), streamSettings) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { conn.Close() }) + return conn.(*transmasque.Conn) + } + h3 := dial() + h2 := dial(http2.NextProtoTLS) + + for _, c := range []struct{ from, to *transmasque.Conn }{{h3, h2}, {h2, h3}} { + for i := range c.from.LocalAddrs() { + src, dst := c.from.LocalAddrs()[i], c.to.LocalAddrs()[i] + payload := make([]byte, 1000) + rand.Read(payload) + if _, err := c.from.Write(masqueIPPacket(src, dst, payload)); err != nil { + t.Fatal(err) + } + received := make(chan []byte, 1) + go func() { + b := make([]byte, 2048) + n, _ := c.to.Read(b) + received <- b[:n] + }() + select { + case p := <-received: + gotSrc, gotDst := masqueIPAddrs(p) + if gotSrc != src || gotDst != dst || !bytes.HasSuffix(p, payload) { + t.Fatalf("unexpected packet from %s to %s: %x", gotSrc, gotDst, p) + } + case <-time.After(5 * time.Second): + t.Fatalf("no packet from %s to %s", src, dst) + } + } + } +} diff --git a/testing/scenarios/shadowsocks_2022_test.go b/testing/scenarios/shadowsocks_2022_test.go index c5e0f8b3d..f456c5a3c 100644 --- a/testing/scenarios/shadowsocks_2022_test.go +++ b/testing/scenarios/shadowsocks_2022_test.go @@ -6,7 +6,6 @@ import ( "testing" "time" - "github.com/sagernet/sing-shadowsocks/shadowaead_2022" "github.com/xtls/xray-core/app/log" "github.com/xtls/xray-core/app/proxyman" "github.com/xtls/xray-core/common" @@ -22,9 +21,19 @@ import ( "golang.org/x/sync/errgroup" ) +var ss2022Methods = []string{ + shadowsocks_2022.MethodAES128GCM, + shadowsocks_2022.MethodAES256GCM, + shadowsocks_2022.MethodChaCha20Poly1305, +} + func TestShadowsocks2022Tcp(t *testing.T) { - for _, method := range shadowaead_2022.List { - password := make([]byte, 32) + for _, method := range ss2022Methods { + keySize := 32 + if method == shadowsocks_2022.MethodAES128GCM { + keySize = 16 + } + password := make([]byte, keySize) rand.Read(password) t.Run(method, func(t *testing.T) { testShadowsocks2022Tcp(t, method, base64.StdEncoding.EncodeToString(password)) @@ -33,21 +42,21 @@ func TestShadowsocks2022Tcp(t *testing.T) { } func TestShadowsocks2022UdpAES128(t *testing.T) { - password := make([]byte, 32) + password := make([]byte, 16) rand.Read(password) - testShadowsocks2022Udp(t, shadowaead_2022.List[0], base64.StdEncoding.EncodeToString(password)) + testShadowsocks2022Udp(t, shadowsocks_2022.MethodAES128GCM, base64.StdEncoding.EncodeToString(password)) } func TestShadowsocks2022UdpAES256(t *testing.T) { password := make([]byte, 32) rand.Read(password) - testShadowsocks2022Udp(t, shadowaead_2022.List[1], base64.StdEncoding.EncodeToString(password)) + testShadowsocks2022Udp(t, shadowsocks_2022.MethodAES256GCM, base64.StdEncoding.EncodeToString(password)) } func TestShadowsocks2022UdpChacha(t *testing.T) { password := make([]byte, 32) rand.Read(password) - testShadowsocks2022Udp(t, shadowaead_2022.List[2], base64.StdEncoding.EncodeToString(password)) + testShadowsocks2022Udp(t, shadowsocks_2022.MethodChaCha20Poly1305, base64.StdEncoding.EncodeToString(password)) } func testShadowsocks2022Tcp(t *testing.T, method string, password string) { diff --git a/transport/internet/config.go b/transport/internet/config.go index 5762655b2..9edd2fbf7 100644 --- a/transport/internet/config.go +++ b/transport/internet/config.go @@ -27,7 +27,7 @@ var strategy = [11][3]byte{ func RegisterProtocolConfigCreator(name string, creator ConfigCreator) error { if _, found := globalTransportConfigCreatorCache[name]; found { - return errors.New("protocol ", name, " is already registered").AtError() + return errors.New("protocol ", name, " is already registered") } globalTransportConfigCreatorCache[name] = creator return nil diff --git a/transport/internet/dialer.go b/transport/internet/dialer.go index 475d0ff95..6d7b905db 100644 --- a/transport/internet/dialer.go +++ b/transport/internet/dialer.go @@ -38,7 +38,7 @@ var transportDialerCache = make(map[string]dialFunc) // RegisterTransportDialer registers a Dialer with given name. func RegisterTransportDialer(protocol string, dialer dialFunc) error { if _, found := transportDialerCache[protocol]; found { - return errors.New(protocol, " dialer already registered").AtError() + return errors.New(protocol, " dialer already registered") } transportDialerCache[protocol] = dialer return nil @@ -58,7 +58,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *MemoryStrea protocol := streamSettings.ProtocolName dialer := transportDialerCache[protocol] if dialer == nil { - return nil, errors.New(protocol, " dialer not registered").AtError() + return nil, errors.New(protocol, " dialer not registered") } return dialer(ctx, dest, streamSettings) } @@ -66,7 +66,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *MemoryStrea if dest.Network == net.Network_UDP { udpDialer := transportDialerCache["udp"] if udpDialer == nil { - return nil, errors.New("UDP dialer not registered").AtError() + return nil, errors.New("UDP dialer not registered") } return udpDialer(ctx, dest, streamSettings) } @@ -86,7 +86,7 @@ var ( func LookupForIP(domain string, strategy DomainStrategy, localAddr net.Address) ([]net.IP, error) { if dnsClient == nil { - return nil, errors.New("DNS client not initialized").AtError() + return nil, errors.New("DNS client not initialized") } ips, _, err := dnsClient.LookupIP(domain, dns.IPOption{ @@ -269,11 +269,11 @@ func DialSystem(ctx context.Context, dest net.Destination, sockopt *SocketConfig if len(sockopt.DialerProxy) > 0 { if obm == nil { - return nil, errors.New("there is no outbound manager for dialerProxy").AtError() + return nil, errors.New("there is no outbound manager for dialerProxy") } h := obm.GetHandler(sockopt.DialerProxy) if h == nil { - return nil, errors.New("there is no outbound handler for dialerProxy").AtError() + return nil, errors.New("there is no outbound handler for dialerProxy") } return redirect(ctx, dest, sockopt.DialerProxy, h), nil } diff --git a/transport/internet/dns_skip.go b/transport/internet/dns_skip.go new file mode 100644 index 000000000..f7be4744e --- /dev/null +++ b/transport/internet/dns_skip.go @@ -0,0 +1,32 @@ +package internet + +import ( + "net/netip" + "slices" + "sync/atomic" +) + +var skippedDNSServers atomic.Pointer[[]netip.Addr] + +// SkipDNSServers has the queries Xray sends to the system's DNS servers on its +// own, like those of localdns, skip servers until it is called again. The DNS +// servers of a TUN are only meant for what goes through it: queried by Xray +// itself they lead back into it, or nowhere. +func SkipDNSServers(servers []netip.Addr) { + skipped := make([]netip.Addr, len(servers)) + for i, server := range servers { + skipped[i] = server.Unmap() + } + skippedDNSServers.Store(&skipped) +} + +// IsSkippedDNSServer reports whether address, a DNS server as host:port, is to +// be skipped, see SkipDNSServers. +func IsSkippedDNSServer(address string) bool { + skipped := skippedDNSServers.Load() + if skipped == nil { + return false + } + server, err := netip.ParseAddrPort(address) + return err == nil && slices.Contains(*skipped, server.Addr().Unmap()) +} diff --git a/transport/internet/dns_skip_test.go b/transport/internet/dns_skip_test.go new file mode 100644 index 000000000..48591b081 --- /dev/null +++ b/transport/internet/dns_skip_test.go @@ -0,0 +1,27 @@ +package internet_test + +import ( + "net/netip" + "testing" + + "github.com/xtls/xray-core/transport/internet" +) + +func TestSkipDNSServers(t *testing.T) { + internet.SkipDNSServers([]netip.Addr{netip.MustParseAddr("::ffff:203.0.113.53"), netip.MustParseAddr("2001:db8::53")}) + t.Cleanup(func() { internet.SkipDNSServers(nil) }) + for address, want := range map[string]bool{ + "203.0.113.53:53": true, + "[2001:db8::53]:53": true, + "198.51.100.53:53": false, + "localhost:53": false, + } { + if got := internet.IsSkippedDNSServer(address); got != want { + t.Errorf("IsSkippedDNSServer(%q) = %v, want %v", address, got, want) + } + } + internet.SkipDNSServers(nil) + if internet.IsSkippedDNSServer("203.0.113.53:53") { + t.Error("still skipped after SkipDNSServers(nil)") + } +} diff --git a/transport/internet/finalmask/finalmask.go b/transport/internet/finalmask/finalmask.go index 179f958f0..a64d4183a 100644 --- a/transport/internet/finalmask/finalmask.go +++ b/transport/internet/finalmask/finalmask.go @@ -5,6 +5,7 @@ import ( "fmt" "slices" + "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" @@ -81,14 +82,14 @@ func (fm *FinalMask) DialTCP(ctx context.Context, dest net.Destination) (net.Con if err != nil { return nil, err } - return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err + return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err }, } for i := range fm.tcpMasks { var newConn net.Conn newConn, err = fm.tcpMasks[i].WrapConnClient(conn, &dest, dialer) if err != nil { - _ = conn.Close() + common.CloseIfExists(conn) return nil, err } conn = newConn @@ -143,7 +144,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con if err != nil { return nil, err } - return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil + return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil } for i := range fm.udpMasks { if i > 0 { @@ -170,7 +171,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con if err != nil { return nil, err } - return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, err + return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, err }, } var sizes []int @@ -193,7 +194,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con } newConn, err = fm.udpMasks[i].WrapPacketConnClient(conn, &dest, dialer) if err != nil { - _ = conn.Close() + common.CloseIfExists(conn) return nil, err } conn = newConn @@ -207,7 +208,7 @@ func (fm *FinalMask) DialUDP(ctx context.Context, dest net.Destination) (net.Con if addr == nil { addr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}} } - return &PacketConnWrapper{PacketConn: conn, udpAddr: addr}, nil + return &net.PacketConnWrapper{PacketConn: conn, Dest: addr}, nil } func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.PacketConn, error) { @@ -240,7 +241,7 @@ func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.Packe if _, ok := fm.udpMasks[i].(interface{ HeaderConn() }); ok { newConn, err = fm.udpMasks[i].WrapPacketConnServer(nil, nil, nil) if err != nil { - _ = conn.Close() + common.CloseIfExists(conn) return nil, err } sizes = append(sizes, newConn.(interface{ Size() int }).Size()) @@ -253,7 +254,7 @@ func (fm *FinalMask) ListenPacket(ctx context.Context, addr net.Addr) (net.Packe } newConn, err = fm.udpMasks[i].WrapPacketConnServer(conn, addr, lc) if err != nil { - _ = conn.Close() + common.CloseIfExists(conn) return nil, err } conn = newConn @@ -271,24 +272,6 @@ const ( UDPSize = 4096 ) -type PacketConnWrapper struct { - net.PacketConn - udpAddr net.Addr -} - -func (c *PacketConnWrapper) RemoteAddr() net.Addr { - return c.udpAddr -} - -func (c *PacketConnWrapper) Read(b []byte) (n int, err error) { - n, _, err = c.PacketConn.ReadFrom(b) - return -} - -func (c *PacketConnWrapper) Write(b []byte) (n int, err error) { - return c.PacketConn.WriteTo(b, c.udpAddr) -} - type headerManagerConn struct { net.PacketConn diff --git a/transport/internet/finalmask/noise/config.pb.go b/transport/internet/finalmask/noise/config.pb.go index 71ba461a6..6f3a5ef16 100644 --- a/transport/internet/finalmask/noise/config.pb.go +++ b/transport/internet/finalmask/noise/config.pb.go @@ -21,6 +21,135 @@ const ( _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) ) +type Segment_Kind int32 + +const ( + Segment_BYTES Segment_Kind = 0 + Segment_RANDOM Segment_Kind = 1 + Segment_RANDOM_ASCII Segment_Kind = 2 + Segment_RANDOM_DIGIT Segment_Kind = 3 + Segment_TIMESTAMP Segment_Kind = 4 + Segment_COUNTER Segment_Kind = 5 + Segment_NONCE Segment_Kind = 6 +) + +// Enum value maps for Segment_Kind. +var ( + Segment_Kind_name = map[int32]string{ + 0: "BYTES", + 1: "RANDOM", + 2: "RANDOM_ASCII", + 3: "RANDOM_DIGIT", + 4: "TIMESTAMP", + 5: "COUNTER", + 6: "NONCE", + } + Segment_Kind_value = map[string]int32{ + "BYTES": 0, + "RANDOM": 1, + "RANDOM_ASCII": 2, + "RANDOM_DIGIT": 3, + "TIMESTAMP": 4, + "COUNTER": 5, + "NONCE": 6, + } +) + +func (x Segment_Kind) Enum() *Segment_Kind { + p := new(Segment_Kind) + *p = x + return p +} + +func (x Segment_Kind) String() string { + return protoimpl.X.EnumStringOf(x.Descriptor(), protoreflect.EnumNumber(x)) +} + +func (Segment_Kind) Descriptor() protoreflect.EnumDescriptor { + return file_transport_internet_finalmask_noise_config_proto_enumTypes[0].Descriptor() +} + +func (Segment_Kind) Type() protoreflect.EnumType { + return &file_transport_internet_finalmask_noise_config_proto_enumTypes[0] +} + +func (x Segment_Kind) Number() protoreflect.EnumNumber { + return protoreflect.EnumNumber(x) +} + +// Deprecated: Use Segment_Kind.Descriptor instead. +func (Segment_Kind) EnumDescriptor() ([]byte, []int) { + return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0, 0} +} + +type Segment struct { + state protoimpl.MessageState `protogen:"open.v1"` + Kind Segment_Kind `protobuf:"varint,1,opt,name=kind,proto3,enum=xray.transport.internet.finalmask.noise.Segment_Kind" json:"kind,omitempty"` + Bytes []byte `protobuf:"bytes,2,opt,name=bytes,proto3" json:"bytes,omitempty"` + MinSize int64 `protobuf:"varint,3,opt,name=min_size,json=minSize,proto3" json:"min_size,omitempty"` + MaxSize int64 `protobuf:"varint,4,opt,name=max_size,json=maxSize,proto3" json:"max_size,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *Segment) Reset() { + *x = Segment{} + mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *Segment) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*Segment) ProtoMessage() {} + +func (x *Segment) ProtoReflect() protoreflect.Message { + mi := &file_transport_internet_finalmask_noise_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 Segment.ProtoReflect.Descriptor instead. +func (*Segment) Descriptor() ([]byte, []int) { + return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0} +} + +func (x *Segment) GetKind() Segment_Kind { + if x != nil { + return x.Kind + } + return Segment_BYTES +} + +func (x *Segment) GetBytes() []byte { + if x != nil { + return x.Bytes + } + return nil +} + +func (x *Segment) GetMinSize() int64 { + if x != nil { + return x.MinSize + } + return 0 +} + +func (x *Segment) GetMaxSize() int64 { + if x != nil { + return x.MaxSize + } + return 0 +} + type Item struct { state protoimpl.MessageState `protogen:"open.v1"` RandMin int64 `protobuf:"varint,1,opt,name=rand_min,json=randMin,proto3" json:"rand_min,omitempty"` @@ -30,13 +159,14 @@ type Item struct { Packet []byte `protobuf:"bytes,5,opt,name=packet,proto3" json:"packet,omitempty"` DelayMin int64 `protobuf:"varint,6,opt,name=delay_min,json=delayMin,proto3" json:"delay_min,omitempty"` DelayMax int64 `protobuf:"varint,7,opt,name=delay_max,json=delayMax,proto3" json:"delay_max,omitempty"` + Segments []*Segment `protobuf:"bytes,8,rep,name=segments,proto3" json:"segments,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } func (x *Item) Reset() { *x = Item{} - mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0] + mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -48,7 +178,7 @@ func (x *Item) String() string { func (*Item) ProtoMessage() {} func (x *Item) ProtoReflect() protoreflect.Message { - mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[0] + mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -61,7 +191,7 @@ func (x *Item) ProtoReflect() protoreflect.Message { // Deprecated: Use Item.ProtoReflect.Descriptor instead. func (*Item) Descriptor() ([]byte, []int) { - return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{0} + return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{1} } func (x *Item) GetRandMin() int64 { @@ -113,6 +243,13 @@ func (x *Item) GetDelayMax() int64 { return 0 } +func (x *Item) GetSegments() []*Segment { + if x != nil { + return x.Segments + } + return nil +} + type Config struct { state protoimpl.MessageState `protogen:"open.v1"` ResetMin int64 `protobuf:"varint,1,opt,name=reset_min,json=resetMin,proto3" json:"reset_min,omitempty"` @@ -124,7 +261,7 @@ type Config struct { func (x *Config) Reset() { *x = Config{} - mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1] + mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[2] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -136,7 +273,7 @@ func (x *Config) String() string { func (*Config) ProtoMessage() {} func (x *Config) ProtoReflect() protoreflect.Message { - mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[1] + mi := &file_transport_internet_finalmask_noise_config_proto_msgTypes[2] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -149,7 +286,7 @@ func (x *Config) ProtoReflect() protoreflect.Message { // Deprecated: Use Config.ProtoReflect.Descriptor instead. func (*Config) Descriptor() ([]byte, []int) { - return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{1} + return file_transport_internet_finalmask_noise_config_proto_rawDescGZIP(), []int{2} } func (x *Config) GetResetMin() int64 { @@ -177,7 +314,21 @@ var File_transport_internet_finalmask_noise_config_proto protoreflect.FileDescri const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" + "\n" + - "/transport/internet/finalmask/noise/config.proto\x12'xray.transport.internet.finalmask.noise\"\xda\x01\n" + + "/transport/internet/finalmask/noise/config.proto\x12'xray.transport.internet.finalmask.noise\"\x8a\x02\n" + + "\aSegment\x12I\n" + + "\x04kind\x18\x01 \x01(\x0e25.xray.transport.internet.finalmask.noise.Segment.KindR\x04kind\x12\x14\n" + + "\x05bytes\x18\x02 \x01(\fR\x05bytes\x12\x19\n" + + "\bmin_size\x18\x03 \x01(\x03R\aminSize\x12\x19\n" + + "\bmax_size\x18\x04 \x01(\x03R\amaxSize\"h\n" + + "\x04Kind\x12\t\n" + + "\x05BYTES\x10\x00\x12\n" + + "\n" + + "\x06RANDOM\x10\x01\x12\x10\n" + + "\fRANDOM_ASCII\x10\x02\x12\x10\n" + + "\fRANDOM_DIGIT\x10\x03\x12\r\n" + + "\tTIMESTAMP\x10\x04\x12\v\n" + + "\aCOUNTER\x10\x05\x12\t\n" + + "\x05NONCE\x10\x06\"\xa8\x02\n" + "\x04Item\x12\x19\n" + "\brand_min\x18\x01 \x01(\x03R\arandMin\x12\x19\n" + "\brand_max\x18\x02 \x01(\x03R\arandMax\x12$\n" + @@ -185,7 +336,8 @@ const file_transport_internet_finalmask_noise_config_proto_rawDesc = "" + "\x0erand_range_max\x18\x04 \x01(\x05R\frandRangeMax\x12\x16\n" + "\x06packet\x18\x05 \x01(\fR\x06packet\x12\x1b\n" + "\tdelay_min\x18\x06 \x01(\x03R\bdelayMin\x12\x1b\n" + - "\tdelay_max\x18\a \x01(\x03R\bdelayMax\"\x87\x01\n" + + "\tdelay_max\x18\a \x01(\x03R\bdelayMax\x12L\n" + + "\bsegments\x18\b \x03(\v20.xray.transport.internet.finalmask.noise.SegmentR\bsegments\"\x87\x01\n" + "\x06Config\x12\x1b\n" + "\treset_min\x18\x01 \x01(\x03R\bresetMin\x12\x1b\n" + "\treset_max\x18\x02 \x01(\x03R\bresetMax\x12C\n" + @@ -204,18 +356,23 @@ func file_transport_internet_finalmask_noise_config_proto_rawDescGZIP() []byte { return file_transport_internet_finalmask_noise_config_proto_rawDescData } -var file_transport_internet_finalmask_noise_config_proto_msgTypes = make([]protoimpl.MessageInfo, 2) +var file_transport_internet_finalmask_noise_config_proto_enumTypes = make([]protoimpl.EnumInfo, 1) +var file_transport_internet_finalmask_noise_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3) var file_transport_internet_finalmask_noise_config_proto_goTypes = []any{ - (*Item)(nil), // 0: xray.transport.internet.finalmask.noise.Item - (*Config)(nil), // 1: xray.transport.internet.finalmask.noise.Config + (Segment_Kind)(0), // 0: xray.transport.internet.finalmask.noise.Segment.Kind + (*Segment)(nil), // 1: xray.transport.internet.finalmask.noise.Segment + (*Item)(nil), // 2: xray.transport.internet.finalmask.noise.Item + (*Config)(nil), // 3: xray.transport.internet.finalmask.noise.Config } var file_transport_internet_finalmask_noise_config_proto_depIdxs = []int32{ - 0, // 0: xray.transport.internet.finalmask.noise.Config.items:type_name -> xray.transport.internet.finalmask.noise.Item - 1, // [1:1] is the sub-list for method output_type - 1, // [1:1] is the sub-list for method input_type - 1, // [1:1] is the sub-list for extension type_name - 1, // [1:1] is the sub-list for extension extendee - 0, // [0:1] is the sub-list for field type_name + 0, // 0: xray.transport.internet.finalmask.noise.Segment.kind:type_name -> xray.transport.internet.finalmask.noise.Segment.Kind + 1, // 1: xray.transport.internet.finalmask.noise.Item.segments:type_name -> xray.transport.internet.finalmask.noise.Segment + 2, // 2: xray.transport.internet.finalmask.noise.Config.items:type_name -> xray.transport.internet.finalmask.noise.Item + 3, // [3:3] is the sub-list for method output_type + 3, // [3:3] is the sub-list for method input_type + 3, // [3:3] is the sub-list for extension type_name + 3, // [3:3] is the sub-list for extension extendee + 0, // [0:3] is the sub-list for field type_name } func init() { file_transport_internet_finalmask_noise_config_proto_init() } @@ -228,13 +385,14 @@ func file_transport_internet_finalmask_noise_config_proto_init() { File: protoimpl.DescBuilder{ GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_noise_config_proto_rawDesc), len(file_transport_internet_finalmask_noise_config_proto_rawDesc)), - NumEnums: 0, - NumMessages: 2, + NumEnums: 1, + NumMessages: 3, NumExtensions: 0, NumServices: 0, }, GoTypes: file_transport_internet_finalmask_noise_config_proto_goTypes, DependencyIndexes: file_transport_internet_finalmask_noise_config_proto_depIdxs, + EnumInfos: file_transport_internet_finalmask_noise_config_proto_enumTypes, MessageInfos: file_transport_internet_finalmask_noise_config_proto_msgTypes, }.Build() File_transport_internet_finalmask_noise_config_proto = out.File diff --git a/transport/internet/finalmask/noise/config.proto b/transport/internet/finalmask/noise/config.proto index d874b973f..319d8dee7 100644 --- a/transport/internet/finalmask/noise/config.proto +++ b/transport/internet/finalmask/noise/config.proto @@ -6,6 +6,22 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/nois option java_package = "com.xray.transport.internet.finalmask.noise"; option java_multiple_files = true; +message Segment { + enum Kind { + BYTES = 0; + RANDOM = 1; + RANDOM_ASCII = 2; + RANDOM_DIGIT = 3; + TIMESTAMP = 4; + COUNTER = 5; + NONCE = 6; + } + Kind kind = 1; + bytes bytes = 2; + int64 min_size = 3; + int64 max_size = 4; +} + message Item { int64 rand_min = 1; int64 rand_max = 2; @@ -14,6 +30,7 @@ message Item { bytes packet = 5; int64 delay_min = 6; int64 delay_max = 7; + repeated Segment segments = 8; } message Config { diff --git a/transport/internet/finalmask/noise/conn.go b/transport/internet/finalmask/noise/conn.go index 8bb115bb4..951071c0e 100644 --- a/transport/internet/finalmask/noise/conn.go +++ b/transport/internet/finalmask/noise/conn.go @@ -1,18 +1,25 @@ package noise import ( + "crypto/rand" + "encoding/binary" "net" "sync" + "sync/atomic" "time" + "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/crypto" ) +const asciiLetters = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ" + type noiseConn struct { net.PacketConn - config *Config - m map[string]time.Time - mu sync.Mutex + config *Config + m map[string]time.Time + mu sync.Mutex + counter atomic.Uint32 } func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) { @@ -27,6 +34,62 @@ func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { return NewConnClient(c, raw) } +func (c *noiseConn) buildPacket(item *Item) []byte { + if len(item.Segments) == 0 { + if item.RandMax > 0 { + buf := make([]byte, crypto.RandBetween(item.RandMin, item.RandMax)) + crypto.RandBytesBetween(buf, byte(item.RandRangeMin), byte(item.RandRangeMax)) + return buf + } + return item.Packet + } + var out []byte + for _, seg := range item.Segments { + out = append(out, c.buildSegment(seg)...) + } + return out +} + +func (c *noiseConn) buildSegment(seg *Segment) []byte { + switch seg.Kind { + case Segment_BYTES: + return seg.Bytes + case Segment_TIMESTAMP: + b := make([]byte, 4) + binary.BigEndian.PutUint32(b, uint32(time.Now().Unix())) + return b + case Segment_COUNTER: + b := make([]byte, 4) + binary.BigEndian.PutUint32(b, c.counter.Add(1)) + return b + case Segment_NONCE: + b := make([]byte, 8) + common.Must2(rand.Read(b)) + return b + default: + size := crypto.RandBetween(seg.MinSize, seg.MaxSize+1) + if size <= 0 { + return nil + } + buf := make([]byte, size) + switch seg.Kind { + case Segment_RANDOM_ASCII: + common.Must2(rand.Read(buf)) + for i := range buf { + buf[i] = asciiLetters[int(buf[i])%len(asciiLetters)] + } + case Segment_RANDOM_DIGIT: + common.Must2(rand.Read(buf)) + for i := range buf { + buf[i] = '0' + buf[i]%10 + } + default: + common.Must2(rand.Read(buf)) + } + return buf + } +} + func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { c.mu.Lock() defer c.mu.Unlock() @@ -35,13 +98,7 @@ func (c *noiseConn) WriteTo(p []byte, addr net.Addr) (n int, err error) { if t.IsZero() || (c.config.ResetMax > 0 && time.Now().After(t)) { for _, item := range c.config.Items { - if item.RandMax > 0 { - buf := make([]byte, crypto.RandBetween(item.RandMin, item.RandMax)) - crypto.RandBytesBetween(buf, byte(item.RandRangeMin), byte(item.RandRangeMax)) - c.PacketConn.WriteTo(buf, addr) - } else { - c.PacketConn.WriteTo(item.Packet, addr) - } + c.PacketConn.WriteTo(c.buildPacket(item), addr) time.Sleep(time.Duration(crypto.RandBetween(item.DelayMin, item.DelayMax)) * time.Millisecond) } } diff --git a/transport/internet/finalmask/noise/conn_test.go b/transport/internet/finalmask/noise/conn_test.go new file mode 100644 index 000000000..5dade71b0 --- /dev/null +++ b/transport/internet/finalmask/noise/conn_test.go @@ -0,0 +1,137 @@ +package noise + +import ( + "bytes" + "encoding/binary" + "net" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +type fakePacketConn struct { + mu sync.Mutex + written [][]byte +} + +func (c *fakePacketConn) WriteTo(p []byte, _ net.Addr) (int, error) { + c.mu.Lock() + defer c.mu.Unlock() + c.written = append(c.written, bytes.Clone(p)) + return len(p), nil +} + +func (c *fakePacketConn) packets() [][]byte { + c.mu.Lock() + defer c.mu.Unlock() + return c.written +} + +func (c *fakePacketConn) ReadFrom(_ []byte) (int, net.Addr, error) { return 0, nil, nil } +func (c *fakePacketConn) Close() error { return nil } +func (c *fakePacketConn) LocalAddr() net.Addr { return &net.UDPAddr{} } +func (c *fakePacketConn) SetDeadline(time.Time) error { return nil } +func (c *fakePacketConn) SetReadDeadline(time.Time) error { return nil } +func (c *fakePacketConn) SetWriteDeadline(time.Time) error { return nil } + +func newConn() *noiseConn { + return &noiseConn{PacketConn: &fakePacketConn{}, config: &Config{}, m: make(map[string]time.Time)} +} + +func TestBuildSegmentBytes(t *testing.T) { + c := newConn() + got := c.buildSegment(&Segment{Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}}) + require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, got) +} + +func TestBuildSegmentTimestamp(t *testing.T) { + c := newConn() + before := time.Now().Unix() + got := c.buildSegment(&Segment{Kind: Segment_TIMESTAMP}) + require.Len(t, got, 4) + ts := int64(binary.BigEndian.Uint32(got)) + require.GreaterOrEqual(t, ts, before) + require.LessOrEqual(t, ts, time.Now().Unix()) +} + +func TestBuildSegmentCounter(t *testing.T) { + c := newConn() + first := binary.BigEndian.Uint32(c.buildSegment(&Segment{Kind: Segment_COUNTER})) + second := binary.BigEndian.Uint32(c.buildSegment(&Segment{Kind: Segment_COUNTER})) + require.Equal(t, uint32(1), first) + require.Equal(t, uint32(2), second) +} + +func TestBuildSegmentNonce(t *testing.T) { + c := newConn() + a := c.buildSegment(&Segment{Kind: Segment_NONCE}) + b := c.buildSegment(&Segment{Kind: Segment_NONCE}) + require.Len(t, a, 8) + require.Len(t, b, 8) + require.NotEqual(t, a, b) +} + +func TestBuildSegmentRandomSizes(t *testing.T) { + c := newConn() + for range 200 { + require.Len(t, c.buildSegment(&Segment{Kind: Segment_RANDOM, MinSize: 24, MaxSize: 24}), 24) + + n := len(c.buildSegment(&Segment{Kind: Segment_RANDOM, MinSize: 20, MaxSize: 32})) + require.GreaterOrEqual(t, n, 20) + require.LessOrEqual(t, n, 32) + + for _, b := range c.buildSegment(&Segment{Kind: Segment_RANDOM_ASCII, MinSize: 40, MaxSize: 40}) { + require.True(t, (b >= 'a' && b <= 'z') || (b >= 'A' && b <= 'Z'), "not a letter: %q", b) + } + for _, b := range c.buildSegment(&Segment{Kind: Segment_RANDOM_DIGIT, MinSize: 40, MaxSize: 40}) { + require.True(t, b >= '0' && b <= '9', "not a digit: %q", b) + } + } +} + +func TestBuildPacketComposite(t *testing.T) { + c := newConn() + item := &Item{Segments: []*Segment{ + {Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}}, + {Kind: Segment_TIMESTAMP}, + {Kind: Segment_RANDOM, MinSize: 24, MaxSize: 24}, + }} + got := c.buildPacket(item) + require.Len(t, got, 4+4+24) + require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, got[:4]) +} + +func TestBuildPacketLegacy(t *testing.T) { + c := newConn() + require.Equal(t, []byte{1, 2, 3}, c.buildPacket(&Item{Packet: []byte{1, 2, 3}})) + require.Len(t, c.buildPacket(&Item{RandMin: 16, RandMax: 17}), 16) +} + +func TestWriteToSendsNoiseThenPayload(t *testing.T) { + raw := &fakePacketConn{} + c := &noiseConn{ + PacketConn: raw, + m: make(map[string]time.Time), + config: &Config{Items: []*Item{ + {Segments: []*Segment{{Kind: Segment_BYTES, Bytes: []byte{0x0d, 0x0a, 0x0d, 0x0a}}, {Kind: Segment_RANDOM, MinSize: 8, MaxSize: 8}}}, + {RandMin: 40, RandMax: 41}, + }}, + } + addr := &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 51820} + payload := []byte("real-handshake") + _, err := c.WriteTo(payload, addr) + require.NoError(t, err) + + sent := raw.packets() + require.Len(t, sent, 3) + require.Len(t, sent[0], 12) + require.Equal(t, []byte{0x0d, 0x0a, 0x0d, 0x0a}, sent[0][:4]) + require.Len(t, sent[1], 40) + require.Equal(t, payload, sent[2]) + + _, err = c.WriteTo(payload, addr) + require.NoError(t, err) + require.Len(t, raw.packets(), 4) +} diff --git a/transport/internet/finalmask/udp_test.go b/transport/internet/finalmask/udp_test.go index 7d5ae5e39..00751da42 100644 --- a/transport/internet/finalmask/udp_test.go +++ b/transport/internet/finalmask/udp_test.go @@ -380,7 +380,7 @@ func TestPacketConnReadWrite(t *testing.T) { t.Fatal(err) } t.Cleanup(func() { clientConn.Close() }) - client := clientConn.(*finalmask.PacketConnWrapper).PacketConn + client := clientConn.(*net.PacketConnWrapper).PacketConn _ = client.SetDeadline(time.Now().Add(time.Second)) _ = server.SetDeadline(time.Now().Add(time.Second)) diff --git a/transport/internet/finalmask/udphop/conn.go b/transport/internet/finalmask/udphop/conn.go index a71a40439..f7bb63d3c 100644 --- a/transport/internet/finalmask/udphop/conn.go +++ b/transport/internet/finalmask/udphop/conn.go @@ -73,7 +73,7 @@ func NewUDPHopConn(c *Config, dest *net.Destination, dialer *finalmask.Dialer) ( if err != nil { return nil, err } - cur := conn.(*finalmask.PacketConnWrapper).PacketConn + cur := conn.(*net.PacketConnWrapper).PacketConn addr := conn.RemoteAddr().(*net.UDPAddr) client := &udpHopConn{ dialer: dialer, @@ -150,7 +150,7 @@ func (c *udpHopConn) hop() { _ = c.pre.Close() } c.pre = c.cur - c.cur = conn.(*finalmask.PacketConnWrapper).PacketConn + c.cur = conn.(*net.PacketConnWrapper).PacketConn c.wg.Add(1) go c.recv(c.cur) } @@ -223,13 +223,6 @@ func (c *udpHopConn) Close() error { } _ = c.cur.Close() c.wg.Wait() - select { - case packet := <-c.readCh: - if packet.p != nil { - pool.Put(packet.p[:cap(packet.p)]) - } - default: - } close(c.readCh) return nil } diff --git a/transport/internet/finalmask/xdns/client.go b/transport/internet/finalmask/xdns/client.go index 557347f21..b976f154f 100644 --- a/transport/internet/finalmask/xdns/client.go +++ b/transport/internet/finalmask/xdns/client.go @@ -1,417 +1,441 @@ package xdns import ( - "bytes" "context" "crypto/rand" - "encoding/base32" - "encoding/binary" - go_errors "errors" "io" - "net" - "strconv" + mrand "math/rand" "sync" "sync/atomic" "time" "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/transport/internet/finalmask" + "golang.org/x/net/dns/dnsmessage" ) const ( - numPadding = 3 - numPaddingForPoll = 8 initPollDelay = 500 * time.Millisecond maxPollDelay = 10 * time.Second pollDelayMultiplier = 2.0 pollLimit = 16 ) -var base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding) +var pool4K = sync.Pool{ + New: func() any { + return make([]byte, 4096) + }, +} type packet struct { p []byte addr net.Addr } -type xdnsConnClient struct { - net.PacketConn +type xdnsClient struct { + dialer *finalmask.Dialer - resolverAddrs []*net.UDPAddr - resolverTypes []uint16 - resolverIdx uint32 - resolverSend map[string]*atomic.Uint32 + clientID ClientID + fragID atomic.Uint32 + domains []*Domain + extraPoll int32 - clientID []byte - domains []Name + resolvers []Resolver + resolverSends []atomic.Uint32 + resolverIndex atomic.Uint32 - pollChan chan struct{} - readQueue chan *packet - writeQueue chan *packet - - closed bool - mutex sync.Mutex + readCh chan packet + sendCh chan []byte + poolCh chan struct{} + closeCh chan struct{} + wg sync.WaitGroup + mu sync.Mutex } -func NewConnClient(c *Config, raw net.PacketConn) (net.PacketConn, error) { +func NewClient(c *Config, dialer *finalmask.Dialer) (net.PacketConn, error) { + if len(c.Domains) == 0 { + return nil, errors.New("empty domains") + } if len(c.Resolvers) == 0 { return nil, errors.New("empty resolvers") } - - var domains []Name - var servers []string - var resolverTypes []uint16 - for _, rs := range c.Resolvers { - domain, server, resolverType, err := parseResolver(rs) - if err != nil { - return nil, errors.New("invalid resolvers").Base(err) - } - domains = append(domains, domain) - servers = append(servers, server) - resolverTypes = append(resolverTypes, resolverType) + if c.ExtraPoll < 0 || c.ExtraPoll > 3 { + return nil, errors.New("c.ExtraPoll < 0 || c.ExtraPoll > 3") } - - var resolverAddrs []*net.UDPAddr - resolverSend := make(map[string]*atomic.Uint32) - for _, rs := range servers { - h, p, err := net.SplitHostPort(rs) + domains := make([]*Domain, 0, len(c.Domains)) + for i := range c.Domains { + 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 := 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 } - ip := net.ParseIP(h) - if ip == nil { - return nil, errors.New("invalid ip address") - } - port, err := strconv.Atoi(p) + domains = append(domains, domain) + } + resolvers := make([]Resolver, 0, len(c.Resolvers)) + for i := range c.Resolvers { + resolver, err := NewResolver(c.Resolvers[i], dialer) if err != nil { - return nil, errors.New("invalid port").Base(err) + return nil, err } - addr := &net.UDPAddr{IP: ip, Port: port} - resolverAddrs = append(resolverAddrs, addr) - resolverSend[addr.String()] = &atomic.Uint32{} + resolvers = append(resolvers, resolver) } + client := &xdnsClient{ + dialer: dialer, - conn := &xdnsConnClient{ - PacketConn: raw, + clientID: NewClientID(), + domains: domains, + extraPoll: c.ExtraPoll, - resolverAddrs: resolverAddrs, - resolverTypes: resolverTypes, - resolverIdx: 0, - resolverSend: resolverSend, + resolvers: resolvers, + resolverSends: make([]atomic.Uint32, len(c.Resolvers)), - clientID: make([]byte, 8), - domains: domains, - - pollChan: make(chan struct{}, pollLimit), - readQueue: make(chan *packet, 256), - writeQueue: make(chan *packet, 256), + readCh: make(chan packet), + sendCh: make(chan []byte, 16), + poolCh: make(chan struct{}, pollLimit), + closeCh: make(chan struct{}), } - - common.Must2(rand.Read(conn.clientID)) - - go conn.recvLoop() - go conn.sendLoop() - - return conn, nil + go client.run() + return client, nil } -func (c *xdnsConnClient) recvLoop() { - var buf [finalmask.UDPSize]byte +func (c *xdnsClient) closed() bool { + select { + case <-c.closeCh: + return true + default: + return false + } +} - for { - if c.closed { +func (c *xdnsClient) read(buf []byte, addr net.Addr) bool { + msg := dnsmessage.Message{} + if err := msg.Unpack(buf); err != nil { + return false + } + if !msg.Header.Response || msg.Header.Truncated || msg.Header.RCode != dnsmessage.RCodeSuccess || len(msg.Questions) != 1 { + return false + } + + var domain *Domain + for i := range c.domains { + if c.domains[i].IsDomain(msg.Questions[0].Name) { + domain = c.domains[i] break } + } + if domain == nil || !domain.HasType(uint16(msg.Questions[0].Type)) { + return false + } - n, addr, err := c.PacketConn.ReadFrom(buf[:]) + edns0 := uint16(0) + for i := range msg.Additionals { + if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT { + edns0 = uint16(msg.Additionals[i].Header.Class) + break + } + } + errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type) + + resp := NewResp(msg, domain, 0) + + p := pool4K.Get().([]byte) + n := resp.Decode(p) + p = p[:n] + + b := p + var bs [][]byte + for len(b) > 1 { + last := b[0]&0xC0 == 0xC0 + length := int(b[0]&0x3F)<<8 | int(b[1]) + b = b[2:] + if length > len(b) { + bs = nil + break + } + packet := make([]byte, length) + copy(packet, b) + bs = append(bs, packet) + if last { + break + } + b = b[length:] + if len(b) < 2 { + bs = nil + } + } + pool4K.Put(p[:cap(p)]) + + for i := range bs { + select { + case <-c.closeCh: + return true + case c.readCh <- packet{p: bs[i], addr: addr}: + } + } + return len(bs) > 0 +} + +func (c *xdnsClient) run() { + for i := range len(c.resolvers) { + c.wg.Add(1) + go c.recv(i) + } + + c.wg.Add(1) + go c.send() + + c.wg.Wait() + close(c.readCh) + close(c.sendCh) + close(c.poolCh) +} + +func (c *xdnsClient) recv(i int) { + defer c.wg.Done() + + var buf [4096]byte + for { + n, err := c.resolvers[i].Read(buf[:]) if err != nil { - if go_errors.Is(err, net.ErrClosed) { - break + if c.closed() { + return } - continue + errors.LogErrorInner(context.Background(), err, "recv err ", i) + return } - - if addr == nil { - continue - } - - send := c.resolverSend[addr.String()] - if send == nil { - continue - } - - resp, err := MessageFromWireFormat(buf[:n]) - if err != nil { - errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err) - continue - } - - payload := dnsResponsePayload(&resp, c.domains) - - r := bytes.NewReader(payload) - anyPacket := false - for { - p, err := nextPacket(r) - if err != nil { - break - } - anyPacket = true - - buf := make([]byte, len(p)) - copy(buf, p) + if c.read(buf[:n], c.resolvers[i].Addr()) { + c.resolverSends[i].Store(0) select { - case c.readQueue <- &packet{ - p: buf, - addr: addr, - }: - default: - errors.LogDebug(context.Background(), addr, " mask read err queue full") - } - } - - if anyPacket { - send.Store(0) - select { - case c.pollChan <- struct{}{}: + case c.poolCh <- struct{}{}: default: } } } - - errors.LogDebug(context.Background(), "xdns closed") - - close(c.pollChan) - close(c.readQueue) - - c.mutex.Lock() - defer c.mutex.Unlock() - - c.closed = true - close(c.writeQueue) } -func (c *xdnsConnClient) sendLoop() { - pollDelay := initPollDelay - pollTimer := time.NewTimer(pollDelay) - for { - var p *packet - pollTimerExpired := false +func (c *xdnsClient) send() { + defer c.wg.Done() - select { - case p = <-c.writeQueue: - default: - select { - case p = <-c.writeQueue: - case <-c.pollChan: - case <-pollTimer.C: - pollTimerExpired = true + var buf [512]byte + var data [255]byte + + sendMsg := func(p []byte, domain *Domain, qtype uint16) { + msg := dnsmessage.Message{ + Header: dnsmessage.Header{ + RecursionDesired: true, + }, + Questions: []dnsmessage.Question{ + { + Name: domain.Encode(p), + Type: dnsmessage.Type(qtype), + Class: dnsmessage.ClassINET, + }, + }, + } + if domain.edns0 > 0 { + msg.Additionals = []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("."), + Type: dnsmessage.TypeOPT, + Class: dnsmessage.Class(domain.edns0), + TTL: 0, + }, + Body: &dnsmessage.OPTResource{}, + }, } } + pack := common.Must2(msg.AppendPack(buf[:0])) + common.Must2(rand.Read(pack[:2])) - if p != nil { - select { - case <-c.pollChan: - default: + index := c.resolverIndex.Load() + cur := c.resolverSends[index].Add(1) + i := index + for { + i++ + if i == uint32(len(c.resolvers)) { + i = 0 } - } else { - encoded, _ := encode(nil, c.clientID, c.domains[c.resolverIdx], c.resolverTypes[c.resolverIdx]) - p = &packet{ - p: encoded, + if i == index { + break + } + if cur > c.resolverSends[i].Load() { + break } } + c.resolverIndex.Store(i) + c.resolvers[index].Send(pack) + } - if pollTimerExpired { - pollDelay = time.Duration(float64(pollDelay) * pollDelayMultiplier) - if pollDelay > maxPollDelay { - pollDelay = maxPollDelay - } - } else { - if !pollTimer.Stop() { - <-pollTimer.C - } - pollDelay = initPollDelay - } - pollTimer.Reset(pollDelay) + send := func(p []byte) { + domain := c.domains[mrand.Intn(len(c.domains))] + qtype := domain.types[mrand.Intn(len(domain.types))] - if c.closed { + if len(p) == 0 { + copy(data[:], c.clientID[:]) + data[0] |= TypeMap[qtype] + data[8] = 8 + common.Must2(rand.Read(data[9:17])) + sendMsg(data[:17], domain, qtype) return } - cur := c.resolverIdx - curSend := c.resolverSend[c.resolverAddrs[cur].String()].Add(1) - _, _ = c.PacketConn.WriteTo(p.p, c.resolverAddrs[cur]) - for { - c.resolverIdx += 1 - c.resolverIdx %= uint32(len(c.resolverAddrs)) - if c.resolverIdx == cur { - break + if len(p) <= domain.cap-12 { + copy(data[:], c.clientID[:]) + data[0] |= TypeMap[qtype] + data[8] = 3 + common.Must2(rand.Read(data[9:12])) + copy(data[12:], p) + sendMsg(data[:12+len(p)], domain, qtype) + return + } + + if len(p) <= 255*(domain.cap-15) { + copy(data[:], c.clientID[:]) + data[0] |= TypeMap[qtype] + data[8] = 3 | 0xC0 + common.Must2(rand.Read(data[9:12])) + + fragID := byte(c.fragID.Add(1)) + fragN := len(p) / (domain.cap - 15) + if len(p)%(domain.cap-15) > 0 { + fragN++ } - if c.resolverSend[c.resolverAddrs[c.resolverIdx].String()].Load() < curSend { - break + + for i := range fragN { + data[12] = fragID + data[13] = byte(i) + data[14] = byte(fragN) + size := min(len(p), domain.cap-15) + copy(data[15:], p[:size]) + sendMsg(data[:15+size], domain, qtype) + p = p[size:] + } + return + } + + errors.LogError(context.Background(), "err size ", len(p)) + } + + ticker := time.NewTicker(initPollDelay) + defer ticker.Stop() + delay := initPollDelay + p := []byte(nil) + timeout := false + for { + select { + case <-c.closeCh: + return + default: + select { + case <-c.closeCh: + return + case p = <-c.sendCh: + case <-c.poolCh: + case <-ticker.C: + timeout = true } } + + if len(p) > 0 { + select { + case <-c.poolCh: + default: + } + } + + send(p) + for range c.extraPoll { + send(nil) + } + + if timeout { + delay *= pollDelayMultiplier + if delay > maxPollDelay { + delay = maxPollDelay + } + timeout = false + } else { + delay = initPollDelay + } + ticker.Reset(delay) } } -func (c *xdnsConnClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) { - packet, ok := <-c.readQueue - if !ok { - return 0, nil, net.ErrClosed +func (c *xdnsClient) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + packet, ok := <-c.readCh + if ok { + return copy(p, packet.p), packet.addr, nil } - if len(p) < len(packet.p) { - errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p)) - return 0, packet.addr, nil - } - copy(p, packet.p) - return len(packet.p), packet.addr, nil + return 0, nil, io.ErrClosedPipe } -func (c *xdnsConnClient) WriteTo(p []byte, addr net.Addr) (n int, err error) { - c.mutex.Lock() - defer c.mutex.Unlock() - - if c.closed { +func (c *xdnsClient) WriteTo(p []byte, addr net.Addr) (n int, err error) { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed() { return 0, io.ErrClosedPipe } - - idx := c.resolverIdx % uint32(len(c.resolverAddrs)) - encoded, err := encode(p, c.clientID, c.domains[idx], c.resolverTypes[idx]) - if err != nil { - errors.LogDebug(context.Background(), addr, " xdns wireformat err ", err, " ", len(p)) - return 0, nil + if len(p) == 0 || len(p) > 4096 { + errors.LogError(context.Background(), "err size ", len(p)) + return 0, errors.New("err size") } - + b := make([]byte, len(p)) + copy(b, p) select { - case c.writeQueue <- &packet{ - p: encoded, - addr: addr, - }: - return len(p), nil + case c.sendCh <- b: default: - errors.LogDebug(context.Background(), addr, " mask write err queue full") - return 0, nil } + return len(p), nil } -func (c *xdnsConnClient) Close() error { - c.closed = true - return c.PacketConn.Close() -} - -func encode(p []byte, clientID []byte, domain Name, qtype uint16) ([]byte, error) { - var decoded []byte - { - if len(p) >= 224 { - return nil, errors.New("too long") - } - var buf bytes.Buffer - buf.Write(clientID[:]) - n := numPadding - if len(p) == 0 { - n = numPaddingForPoll - } - buf.WriteByte(byte(224 + n)) - _, _ = io.CopyN(&buf, rand.Reader, int64(n)) - if len(p) > 0 { - buf.WriteByte(byte(len(p))) - buf.Write(p) - } - decoded = buf.Bytes() - } - - encoded := make([]byte, base32Encoding.EncodedLen(len(decoded))) - base32Encoding.Encode(encoded, decoded) - encoded = bytes.ToLower(encoded) - labels := chunks(encoded, 63) - labels = append(labels, domain...) - name, err := NewName(labels) - if err != nil { - return nil, err - } - - var id uint16 - _ = binary.Read(rand.Reader, binary.BigEndian, &id) - query := &Message{ - ID: id, - Flags: 0x0100, - Question: []Question{ - { - Name: name, - Type: qtype, - Class: ClassIN, - }, - }, - Additional: []RR{ - { - Name: Name{}, - Type: RRTypeOPT, - Class: 4096, - TTL: 0, - Data: []byte{}, - }, - }, - } - - buf, err := query.WireFormat() - if err != nil { - return nil, err - } - - return buf, nil -} - -func chunks(p []byte, n int) [][]byte { - var result [][]byte - for len(p) > 0 { - sz := len(p) - if sz > n { - sz = n - } - result = append(result, p[:sz]) - p = p[sz:] - } - return result -} - -func nextPacket(r *bytes.Reader) ([]byte, error) { - var n uint16 - err := binary.Read(r, binary.BigEndian, &n) - if err != nil { - return nil, err - } - p := make([]byte, n) - _, err = io.ReadFull(r, p) - if err == io.EOF { - err = io.ErrUnexpectedEOF - } - return p, err -} - -func dnsResponsePayload(resp *Message, domains []Name) []byte { - if resp.Flags&0x8000 != 0x8000 { +func (c *xdnsClient) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed() { return nil } - if resp.Flags&0x000f != RcodeNoError { - return nil + close(c.closeCh) + for i := range c.resolvers { + c.resolvers[i].Close() } + return nil +} - if len(resp.Answer) == 0 { - return nil - } - - for _, answer := range resp.Answer { - var ok bool - for _, domain := range domains { - _, ok = answer.Name.TrimSuffix(domain) - if ok { - break - } - } - if !ok { - return nil - } - } - - return decodeResponsePayload(resp.Answer) +func (c *xdnsClient) LocalAddr() net.Addr { return &net.UDPAddr{IP: []byte{0, 0, 0, 0}} } + +func (c *xdnsClient) SetDeadline(t time.Time) error { return errors.New("not support") } + +func (c *xdnsClient) SetReadDeadline(t time.Time) error { return errors.New("not support") } + +func (c *xdnsClient) SetWriteDeadline(t time.Time) error { return errors.New("not support") } + +type ClientID [8]byte + +func NewClientID() ClientID { + var id ClientID + common.Must2(rand.Read(id[:])) + id[0] &= 0xFC + return id +} + +func ClientIDFromRaw(id [8]byte) ClientID { + id[0] &= 0xFC + return id +} + +func ClientIDFromAddr(addr *net.UDPAddr) ClientID { + return ClientID(addr.IP[8:]) +} + +func (id ClientID) Addr() *net.UDPAddr { + var ip [16]byte + ip[0] = 0xFD + copy(ip[8:], id[:]) + return &net.UDPAddr{IP: ip[:]} } diff --git a/transport/internet/finalmask/xdns/config.go b/transport/internet/finalmask/xdns/config.go index 7bae597ab..d982c0aa1 100644 --- a/transport/internet/finalmask/xdns/config.go +++ b/transport/internet/finalmask/xdns/config.go @@ -6,9 +6,9 @@ import ( ) func (c *Config) WrapPacketConnClient(conn net.PacketConn, dest *net.Destination, dialer *finalmask.Dialer) (net.PacketConn, error) { - return NewConnClient(c, conn) + return NewClient(c, dialer) } func (c *Config) WrapPacketConnServer(conn net.PacketConn, addr net.Addr, lc *finalmask.ListenConfig) (net.PacketConn, error) { - return NewConnServer(c, conn) + return NewServer(c, conn) } diff --git a/transport/internet/finalmask/xdns/config.pb.go b/transport/internet/finalmask/xdns/config.pb.go index e1f06aa93..7db693f98 100644 --- a/transport/internet/finalmask/xdns/config.pb.go +++ b/transport/internet/finalmask/xdns/config.pb.go @@ -7,6 +7,7 @@ package xdns import ( + serial "github.com/xtls/xray-core/common/serial" protoreflect "google.golang.org/protobuf/reflect/protoreflect" protoimpl "google.golang.org/protobuf/runtime/protoimpl" reflect "reflect" @@ -21,17 +22,94 @@ const ( _ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20) ) +type DomainProto struct { + state protoimpl.MessageState `protogen:"open.v1"` + Name string `protobuf:"bytes,1,opt,name=name,proto3" json:"name,omitempty"` + LenLimit int32 `protobuf:"varint,2,opt,name=len_limit,json=lenLimit,proto3" json:"len_limit,omitempty"` + LabelLimit int32 `protobuf:"varint,3,opt,name=label_limit,json=labelLimit,proto3" json:"label_limit,omitempty"` + Types []int32 `protobuf:"varint,4,rep,packed,name=types,proto3" json:"types,omitempty"` + Edns0 int32 `protobuf:"varint,5,opt,name=edns0,proto3" json:"edns0,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *DomainProto) Reset() { + *x = DomainProto{} + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *DomainProto) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*DomainProto) ProtoMessage() {} + +func (x *DomainProto) ProtoReflect() protoreflect.Message { + mi := &file_transport_internet_finalmask_xdns_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 DomainProto.ProtoReflect.Descriptor instead. +func (*DomainProto) Descriptor() ([]byte, []int) { + return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0} +} + +func (x *DomainProto) GetName() string { + if x != nil { + return x.Name + } + return "" +} + +func (x *DomainProto) GetLenLimit() int32 { + if x != nil { + return x.LenLimit + } + return 0 +} + +func (x *DomainProto) GetLabelLimit() int32 { + if x != nil { + return x.LabelLimit + } + return 0 +} + +func (x *DomainProto) GetTypes() []int32 { + if x != nil { + return x.Types + } + return nil +} + +func (x *DomainProto) GetEdns0() int32 { + if x != nil { + return x.Edns0 + } + return 0 +} + type Config struct { state protoimpl.MessageState `protogen:"open.v1"` - Domains []string `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"` - Resolvers []string `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"` + Domains []*DomainProto `protobuf:"bytes,1,rep,name=domains,proto3" json:"domains,omitempty"` + Resolvers []*serial.TypedMessage `protobuf:"bytes,2,rep,name=resolvers,proto3" json:"resolvers,omitempty"` + ExtraPoll int32 `protobuf:"varint,3,opt,name=extra_poll,json=extraPoll,proto3" json:"extra_poll,omitempty"` unknownFields protoimpl.UnknownFields sizeCache protoimpl.SizeCache } func (x *Config) Reset() { *x = Config{} - mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1] ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) ms.StoreMessageInfo(mi) } @@ -43,7 +121,7 @@ func (x *Config) String() string { func (*Config) ProtoMessage() {} func (x *Config) ProtoReflect() protoreflect.Message { - mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[0] + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[1] if x != nil { ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) if ms.LoadMessageInfo() == nil { @@ -56,31 +134,139 @@ func (x *Config) ProtoReflect() protoreflect.Message { // Deprecated: Use Config.ProtoReflect.Descriptor instead. func (*Config) Descriptor() ([]byte, []int) { - return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{0} + return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{1} } -func (x *Config) GetDomains() []string { +func (x *Config) GetDomains() []*DomainProto { if x != nil { return x.Domains } return nil } -func (x *Config) GetResolvers() []string { +func (x *Config) GetResolvers() []*serial.TypedMessage { if x != nil { return x.Resolvers } return nil } +func (x *Config) GetExtraPoll() int32 { + if x != nil { + return x.ExtraPoll + } + return 0 +} + +type TCPResolverProto struct { + state protoimpl.MessageState `protogen:"open.v1"` + Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *TCPResolverProto) Reset() { + *x = TCPResolverProto{} + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[2] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *TCPResolverProto) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*TCPResolverProto) ProtoMessage() {} + +func (x *TCPResolverProto) ProtoReflect() protoreflect.Message { + mi := &file_transport_internet_finalmask_xdns_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 TCPResolverProto.ProtoReflect.Descriptor instead. +func (*TCPResolverProto) Descriptor() ([]byte, []int) { + return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{2} +} + +func (x *TCPResolverProto) GetAddr() string { + if x != nil { + return x.Addr + } + return "" +} + +type UDPResolverProto struct { + state protoimpl.MessageState `protogen:"open.v1"` + Addr string `protobuf:"bytes,1,opt,name=addr,proto3" json:"addr,omitempty"` + unknownFields protoimpl.UnknownFields + sizeCache protoimpl.SizeCache +} + +func (x *UDPResolverProto) Reset() { + *x = UDPResolverProto{} + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3] + ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x)) + ms.StoreMessageInfo(mi) +} + +func (x *UDPResolverProto) String() string { + return protoimpl.X.MessageStringOf(x) +} + +func (*UDPResolverProto) ProtoMessage() {} + +func (x *UDPResolverProto) ProtoReflect() protoreflect.Message { + mi := &file_transport_internet_finalmask_xdns_config_proto_msgTypes[3] + 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 UDPResolverProto.ProtoReflect.Descriptor instead. +func (*UDPResolverProto) Descriptor() ([]byte, []int) { + return file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP(), []int{3} +} + +func (x *UDPResolverProto) GetAddr() string { + if x != nil { + return x.Addr + } + return "" +} + var File_transport_internet_finalmask_xdns_config_proto protoreflect.FileDescriptor const file_transport_internet_finalmask_xdns_config_proto_rawDesc = "" + "\n" + - ".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\"@\n" + - "\x06Config\x12\x18\n" + - "\adomains\x18\x01 \x03(\tR\adomains\x12\x1c\n" + - "\tresolvers\x18\x02 \x03(\tR\tresolversB\x94\x01\n" + + ".transport/internet/finalmask/xdns/config.proto\x12&xray.transport.internet.finalmask.xdns\x1a!common/serial/typed_message.proto\"\x8b\x01\n" + + "\vDomainProto\x12\x12\n" + + "\x04name\x18\x01 \x01(\tR\x04name\x12\x1b\n" + + "\tlen_limit\x18\x02 \x01(\x05R\blenLimit\x12\x1f\n" + + "\vlabel_limit\x18\x03 \x01(\x05R\n" + + "labelLimit\x12\x14\n" + + "\x05types\x18\x04 \x03(\x05R\x05types\x12\x14\n" + + "\x05edns0\x18\x05 \x01(\x05R\x05edns0\"\xb6\x01\n" + + "\x06Config\x12M\n" + + "\adomains\x18\x01 \x03(\v23.xray.transport.internet.finalmask.xdns.DomainProtoR\adomains\x12>\n" + + "\tresolvers\x18\x02 \x03(\v2 .xray.common.serial.TypedMessageR\tresolvers\x12\x1d\n" + + "\n" + + "extra_poll\x18\x03 \x01(\x05R\textraPoll\"&\n" + + "\x10TCPResolverProto\x12\x12\n" + + "\x04addr\x18\x01 \x01(\tR\x04addr\"&\n" + + "\x10UDPResolverProto\x12\x12\n" + + "\x04addr\x18\x01 \x01(\tR\x04addrB\x94\x01\n" + "*com.xray.transport.internet.finalmask.xdnsP\x01Z;github.com/xtls/xray-core/transport/internet/finalmask/xdns\xaa\x02&Xray.Transport.Internet.Finalmask.Xdnsb\x06proto3" var ( @@ -95,16 +281,22 @@ func file_transport_internet_finalmask_xdns_config_proto_rawDescGZIP() []byte { return file_transport_internet_finalmask_xdns_config_proto_rawDescData } -var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 1) +var file_transport_internet_finalmask_xdns_config_proto_msgTypes = make([]protoimpl.MessageInfo, 4) var file_transport_internet_finalmask_xdns_config_proto_goTypes = []any{ - (*Config)(nil), // 0: xray.transport.internet.finalmask.xdns.Config + (*DomainProto)(nil), // 0: xray.transport.internet.finalmask.xdns.DomainProto + (*Config)(nil), // 1: xray.transport.internet.finalmask.xdns.Config + (*TCPResolverProto)(nil), // 2: xray.transport.internet.finalmask.xdns.TCPResolverProto + (*UDPResolverProto)(nil), // 3: xray.transport.internet.finalmask.xdns.UDPResolverProto + (*serial.TypedMessage)(nil), // 4: xray.common.serial.TypedMessage } var file_transport_internet_finalmask_xdns_config_proto_depIdxs = []int32{ - 0, // [0:0] is the sub-list for method output_type - 0, // [0:0] is the sub-list for method input_type - 0, // [0:0] is the sub-list for extension type_name - 0, // [0:0] is the sub-list for extension extendee - 0, // [0:0] is the sub-list for field type_name + 0, // 0: xray.transport.internet.finalmask.xdns.Config.domains:type_name -> xray.transport.internet.finalmask.xdns.DomainProto + 4, // 1: xray.transport.internet.finalmask.xdns.Config.resolvers:type_name -> xray.common.serial.TypedMessage + 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_transport_internet_finalmask_xdns_config_proto_init() } @@ -118,7 +310,7 @@ func file_transport_internet_finalmask_xdns_config_proto_init() { GoPackagePath: reflect.TypeOf(x{}).PkgPath(), RawDescriptor: unsafe.Slice(unsafe.StringData(file_transport_internet_finalmask_xdns_config_proto_rawDesc), len(file_transport_internet_finalmask_xdns_config_proto_rawDesc)), NumEnums: 0, - NumMessages: 1, + NumMessages: 4, NumExtensions: 0, NumServices: 0, }, diff --git a/transport/internet/finalmask/xdns/config.proto b/transport/internet/finalmask/xdns/config.proto index b859b17ae..1e464cb2d 100644 --- a/transport/internet/finalmask/xdns/config.proto +++ b/transport/internet/finalmask/xdns/config.proto @@ -6,7 +6,26 @@ option go_package = "github.com/xtls/xray-core/transport/internet/finalmask/xdns option java_package = "com.xray.transport.internet.finalmask.xdns"; option java_multiple_files = true; +import "common/serial/typed_message.proto"; + +message DomainProto { + string name = 1; + int32 len_limit = 2; + int32 label_limit = 3; + repeated int32 types = 4; + int32 edns0 = 5; +} + message Config { - repeated string domains = 1; - repeated string resolvers = 2; + repeated DomainProto domains = 1; + repeated xray.common.serial.TypedMessage resolvers = 2; + int32 extra_poll = 3; +} + +message TCPResolverProto { + string addr = 1; +} + +message UDPResolverProto { + string addr = 1; } \ No newline at end of file diff --git a/transport/internet/finalmask/xdns/dns.go b/transport/internet/finalmask/xdns/dns.go deleted file mode 100644 index 774903857..000000000 --- a/transport/internet/finalmask/xdns/dns.go +++ /dev/null @@ -1,581 +0,0 @@ -// Package dns deals with encoding and decoding DNS wire format. -package xdns - -import ( - "bytes" - "encoding/binary" - "errors" - "fmt" - "io" - "strings" -) - -// The maximum number of DNS name compression pointers we are willing to follow. -// Without something like this, infinite loops are possible. -const compressionPointerLimit = 10 - -var ( - // ErrZeroLengthLabel is the error returned for names that contain a - // zero-length label, like "example..com". - ErrZeroLengthLabel = errors.New("name contains a zero-length label") - - // ErrLabelTooLong is the error returned for labels that are longer than - // 63 octets. - ErrLabelTooLong = errors.New("name contains a label longer than 63 octets") - - // ErrNameTooLong is the error returned for names whose encoded - // representation is longer than 255 octets. - ErrNameTooLong = errors.New("name is longer than 255 octets") - - // ErrReservedLabelType is the error returned when reading a label type - // prefix whose two most significant bits are not 00 or 11. - ErrReservedLabelType = errors.New("reserved label type") - - // ErrTooManyPointers is the error returned when reading a compressed - // name that has too many compression pointers. - ErrTooManyPointers = errors.New("too many compression pointers") - - // ErrTrailingBytes is the error returned when bytes remain in the parse - // buffer after parsing a message. - ErrTrailingBytes = errors.New("trailing bytes after message") - - // ErrIntegerOverflow is the error returned when trying to encode an - // integer greater than 65535 into a 16-bit field. - ErrIntegerOverflow = errors.New("integer overflow") -) - -const ( - // https://tools.ietf.org/html/rfc1035#section-3.2.2 - RRTypeA = 1 - // https://tools.ietf.org/html/rfc1035#section-3.2.2 - RRTypeCNAME = 5 - // https://tools.ietf.org/html/rfc1035#section-3.2.2 - RRTypeTXT = 16 - // https://tools.ietf.org/html/rfc3596#section-2.1 - RRTypeAAAA = 28 - // https://tools.ietf.org/html/rfc6891#section-6.1.1 - RRTypeOPT = 41 - - // https://tools.ietf.org/html/rfc1035#section-3.2.4 - ClassIN = 1 - - // https://tools.ietf.org/html/rfc1035#section-4.1.1 - RcodeNoError = 0 // a.k.a. NOERROR - RcodeFormatError = 1 // a.k.a. FORMERR - RcodeNameError = 3 // a.k.a. NXDOMAIN - RcodeNotImplemented = 4 // a.k.a. NOTIMPL - // https://tools.ietf.org/html/rfc6891#section-9 - ExtendedRcodeBadVers = 16 // a.k.a. BADVERS -) - -// Name represents a domain name, a sequence of labels each of which is 63 -// octets or less in length. -// -// https://tools.ietf.org/html/rfc1035#section-3.1 -type Name [][]byte - -// NewName returns a Name from a slice of labels, after checking the labels for -// validity. Does not include a zero-length label at the end of the slice. -func NewName(labels [][]byte) (Name, error) { - name := Name(labels) - // https://tools.ietf.org/html/rfc1035#section-2.3.4 - // Various objects and parameters in the DNS have size limits. - // labels 63 octets or less - // names 255 octets or less - for _, label := range labels { - if len(label) == 0 { - return nil, ErrZeroLengthLabel - } - if len(label) > 63 { - return nil, ErrLabelTooLong - } - } - // Check the total length. - builder := newMessageBuilder() - builder.WriteName(name) - if len(builder.Bytes()) > 255 { - return nil, ErrNameTooLong - } - return name, nil -} - -// ParseName returns a new Name from a string of labels separated by dots, after -// checking the name for validity. A single dot at the end of the string is -// ignored. -func ParseName(s string) (Name, error) { - b := bytes.TrimSuffix([]byte(s), []byte(".")) - if len(b) == 0 { - // bytes.Split(b, ".") would return [""] in this case - return NewName([][]byte{}) - } else { - return NewName(bytes.Split(b, []byte("."))) - } -} - -// String returns a reversible string representation of name. Labels are -// separated by dots, and any bytes in a label that are outside the set -// [0-9A-Za-z-] are replaced with a \xXX hex escape sequence. -func (name Name) String() string { - if len(name) == 0 { - return "." - } - - var buf strings.Builder - for i, label := range name { - if i > 0 { - buf.WriteByte('.') - } - for _, b := range label { - if b == '-' || - ('0' <= b && b <= '9') || - ('A' <= b && b <= 'Z') || - ('a' <= b && b <= 'z') { - buf.WriteByte(b) - } else { - fmt.Fprintf(&buf, "\\x%02x", b) - } - } - } - return buf.String() -} - -// TrimSuffix returns a Name with the given suffix removed, if it was present. -// The second return value indicates whether the suffix was present. If the -// suffix was not present, the first return value is nil. -func (name Name) TrimSuffix(suffix Name) (Name, bool) { - if len(name) < len(suffix) { - return nil, false - } - split := len(name) - len(suffix) - fore, aft := name[:split], name[split:] - for i := 0; i < len(aft); i++ { - if !bytes.Equal(bytes.ToLower(aft[i]), bytes.ToLower(suffix[i])) { - return nil, false - } - } - return fore, true -} - -// Message represents a DNS message. -// -// https://tools.ietf.org/html/rfc1035#section-4.1 -type Message struct { - ID uint16 - Flags uint16 - - Question []Question - Answer []RR - Authority []RR - Additional []RR -} - -// Opcode extracts the OPCODE part of the Flags field. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.1 -func (message *Message) Opcode() uint16 { - return (message.Flags >> 11) & 0xf -} - -// Rcode extracts the RCODE part of the Flags field. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.1 -func (message *Message) Rcode() uint16 { - return message.Flags & 0x000f -} - -// Question represents an entry in the question section of a message. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.2 -type Question struct { - Name Name - Type uint16 - Class uint16 -} - -// RR represents a resource record. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.3 -type RR struct { - Name Name - Type uint16 - Class uint16 - TTL uint32 - Data []byte -} - -// readName parses a DNS name from r. It leaves r positioned just after the -// parsed name. -func readName(r io.ReadSeeker) (Name, error) { - var labels [][]byte - // We limit the number of compression pointers we are willing to follow. - numPointers := 0 - // If we followed any compression pointers, we must finally seek to just - // past the first pointer. - var seekTo int64 -loop: - for { - var labelType byte - err := binary.Read(r, binary.BigEndian, &labelType) - if err != nil { - return nil, err - } - - switch labelType & 0xc0 { - case 0x00: - // This is an ordinary label. - // https://tools.ietf.org/html/rfc1035#section-3.1 - length := int(labelType & 0x3f) - if length == 0 { - break loop - } - label := make([]byte, length) - _, err := io.ReadFull(r, label) - if err != nil { - return nil, err - } - labels = append(labels, label) - case 0xc0: - // This is a compression pointer. - // https://tools.ietf.org/html/rfc1035#section-4.1.4 - upper := labelType & 0x3f - var lower byte - err := binary.Read(r, binary.BigEndian, &lower) - if err != nil { - return nil, err - } - offset := (uint16(upper) << 8) | uint16(lower) - - if numPointers == 0 { - // The first time we encounter a pointer, - // remember our position so we can seek back to - // it when done. - seekTo, err = r.Seek(0, io.SeekCurrent) - if err != nil { - return nil, err - } - } - numPointers++ - if numPointers > compressionPointerLimit { - return nil, ErrTooManyPointers - } - - // Follow the pointer and continue. - _, err = r.Seek(int64(offset), io.SeekStart) - if err != nil { - return nil, err - } - default: - // "The 10 and 01 combinations are reserved for future - // use." - return nil, ErrReservedLabelType - } - } - // If we followed any pointers, then seek back to just after the first - // one. - if numPointers > 0 { - _, err := r.Seek(seekTo, io.SeekStart) - if err != nil { - return nil, err - } - } - return NewName(labels) -} - -// readQuestion parses one entry from the Question section. It leaves r -// positioned just after the parsed entry. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.2 -func readQuestion(r io.ReadSeeker) (Question, error) { - var question Question - var err error - question.Name, err = readName(r) - if err != nil { - return question, err - } - for _, ptr := range []*uint16{&question.Type, &question.Class} { - err := binary.Read(r, binary.BigEndian, ptr) - if err != nil { - return question, err - } - } - - return question, nil -} - -// readRR parses one resource record. It leaves r positioned just after the -// parsed resource record. -// -// https://tools.ietf.org/html/rfc1035#section-4.1.3 -func readRR(r io.ReadSeeker) (RR, error) { - var rr RR - var err error - rr.Name, err = readName(r) - if err != nil { - return rr, err - } - for _, ptr := range []*uint16{&rr.Type, &rr.Class} { - err := binary.Read(r, binary.BigEndian, ptr) - if err != nil { - return rr, err - } - } - err = binary.Read(r, binary.BigEndian, &rr.TTL) - if err != nil { - return rr, err - } - var rdLength uint16 - err = binary.Read(r, binary.BigEndian, &rdLength) - if err != nil { - return rr, err - } - rr.Data = make([]byte, rdLength) - _, err = io.ReadFull(r, rr.Data) - if err != nil { - return rr, err - } - - return rr, nil -} - -// readMessage parses a complete DNS message. It leaves r positioned just after -// the parsed message. -func readMessage(r io.ReadSeeker) (Message, error) { - var message Message - - // Header section - // https://tools.ietf.org/html/rfc1035#section-4.1.1 - var qdCount, anCount, nsCount, arCount uint16 - for _, ptr := range []*uint16{ - &message.ID, &message.Flags, - &qdCount, &anCount, &nsCount, &arCount, - } { - err := binary.Read(r, binary.BigEndian, ptr) - if err != nil { - return message, err - } - } - - // Question section - // https://tools.ietf.org/html/rfc1035#section-4.1.2 - for i := 0; i < int(qdCount); i++ { - question, err := readQuestion(r) - if err != nil { - return message, err - } - message.Question = append(message.Question, question) - } - - // Answer, Authority, and Additional sections - // https://tools.ietf.org/html/rfc1035#section-4.1.3 - for _, rec := range []struct { - ptr *[]RR - count uint16 - }{ - {&message.Answer, anCount}, - {&message.Authority, nsCount}, - {&message.Additional, arCount}, - } { - for i := 0; i < int(rec.count); i++ { - rr, err := readRR(r) - if err != nil { - return message, err - } - *rec.ptr = append(*rec.ptr, rr) - } - } - - return message, nil -} - -// MessageFromWireFormat parses a message from buf and returns a Message object. -// It returns ErrTrailingBytes if there are bytes remaining in buf after parsing -// is done. -func MessageFromWireFormat(buf []byte) (Message, error) { - r := bytes.NewReader(buf) - message, err := readMessage(r) - if err == io.EOF { - err = io.ErrUnexpectedEOF - } else if err == nil { - // Check for trailing bytes. - _, err = r.ReadByte() - if err == io.EOF { - err = nil - } else if err == nil { - err = ErrTrailingBytes - } - } - return message, err -} - -// messageBuilder manages the state of serializing a DNS message. Its main -// function is to keep track of names already written for the purpose of name -// compression. -type messageBuilder struct { - w bytes.Buffer - nameCache map[string]int -} - -// newMessageBuilder creates a new messageBuilder with an empty name cache. -func newMessageBuilder() *messageBuilder { - return &messageBuilder{ - nameCache: make(map[string]int), - } -} - -// Bytes returns the serialized DNS message as a slice of bytes. -func (builder *messageBuilder) Bytes() []byte { - return builder.w.Bytes() -} - -// WriteName appends name to the in-progress messageBuilder, employing -// compression pointers to previously written names if possible. -func (builder *messageBuilder) WriteName(name Name) { - // https://tools.ietf.org/html/rfc1035#section-3.1 - for i := range name { - // Has this suffix already been encoded in the message? - if ptr, ok := builder.nameCache[name[i:].String()]; ok && ptr&0x3fff == ptr { - // If so, we can write a compression pointer. - binary.Write(&builder.w, binary.BigEndian, uint16(0xc000|ptr)) - return - } - // Not cached; we must encode this label verbatim. Store a cache - // entry pointing to the beginning of it. - builder.nameCache[name[i:].String()] = builder.w.Len() - length := len(name[i]) - if length == 0 || length > 63 { - panic(length) - } - builder.w.WriteByte(byte(length)) - builder.w.Write(name[i]) - } - builder.w.WriteByte(0) -} - -// WriteQuestion appends a Question section entry to the in-progress -// messageBuilder. -func (builder *messageBuilder) WriteQuestion(question *Question) { - // https://tools.ietf.org/html/rfc1035#section-4.1.2 - builder.WriteName(question.Name) - binary.Write(&builder.w, binary.BigEndian, question.Type) - binary.Write(&builder.w, binary.BigEndian, question.Class) -} - -// WriteRR appends a resource record to the in-progress messageBuilder. It -// returns ErrIntegerOverflow if the length of rr.Data does not fit in 16 bits. -func (builder *messageBuilder) WriteRR(rr *RR) error { - // https://tools.ietf.org/html/rfc1035#section-4.1.3 - builder.WriteName(rr.Name) - binary.Write(&builder.w, binary.BigEndian, rr.Type) - binary.Write(&builder.w, binary.BigEndian, rr.Class) - binary.Write(&builder.w, binary.BigEndian, rr.TTL) - rdLength := uint16(len(rr.Data)) - if int(rdLength) != len(rr.Data) { - return ErrIntegerOverflow - } - binary.Write(&builder.w, binary.BigEndian, rdLength) - builder.w.Write(rr.Data) - return nil -} - -// WriteMessage appends a complete DNS message to the in-progress -// messageBuilder. It returns ErrIntegerOverflow if the number of entries in any -// section, or the length of the data in any resource record, does not fit in 16 -// bits. -func (builder *messageBuilder) WriteMessage(message *Message) error { - // Header section - // https://tools.ietf.org/html/rfc1035#section-4.1.1 - binary.Write(&builder.w, binary.BigEndian, message.ID) - binary.Write(&builder.w, binary.BigEndian, message.Flags) - for _, count := range []int{ - len(message.Question), - len(message.Answer), - len(message.Authority), - len(message.Additional), - } { - count16 := uint16(count) - if int(count16) != count { - return ErrIntegerOverflow - } - binary.Write(&builder.w, binary.BigEndian, count16) - } - - // Question section - // https://tools.ietf.org/html/rfc1035#section-4.1.2 - for _, question := range message.Question { - builder.WriteQuestion(&question) - } - - // Answer, Authority, and Additional sections - // https://tools.ietf.org/html/rfc1035#section-4.1.3 - for _, rrs := range [][]RR{message.Answer, message.Authority, message.Additional} { - for _, rr := range rrs { - err := builder.WriteRR(&rr) - if err != nil { - return err - } - } - } - - return nil -} - -// WireFormat encodes a Message as a slice of bytes in DNS wire format. It -// returns ErrIntegerOverflow if the number of entries in any section, or the -// length of the data in any resource record, does not fit in 16 bits. -func (message *Message) WireFormat() ([]byte, error) { - builder := newMessageBuilder() - err := builder.WriteMessage(message) - if err != nil { - return nil, err - } - return builder.Bytes(), nil -} - -// DecodeRDataTXT decodes TXT-DATA (as found in the RDATA for a resource record -// with TYPE=TXT) as a raw byte slice, by concatenating all the -// s it contains. -// -// https://tools.ietf.org/html/rfc1035#section-3.3.14 -func DecodeRDataTXT(p []byte) ([]byte, error) { - var buf bytes.Buffer - for { - if len(p) == 0 { - return nil, io.ErrUnexpectedEOF - } - n := int(p[0]) - p = p[1:] - if len(p) < n { - return nil, io.ErrUnexpectedEOF - } - buf.Write(p[:n]) - p = p[n:] - if len(p) == 0 { - break - } - } - return buf.Bytes(), nil -} - -// EncodeRDataTXT encodes a slice of bytes as TXT-DATA, as appropriate for the -// RDATA of a resource record with TYPE=TXT. No length restriction is enforced -// here; that must be checked at a higher level. -// -// https://tools.ietf.org/html/rfc1035#section-3.3.14 -func EncodeRDataTXT(p []byte) []byte { - // https://tools.ietf.org/html/rfc1035#section-3.3 - // https://tools.ietf.org/html/rfc1035#section-3.3.14 - // TXT data is a sequence of one or more s, where - // is a length octet followed by that number of - // octets. - var buf bytes.Buffer - for len(p) > 255 { - buf.WriteByte(255) - buf.Write(p[:255]) - p = p[255:] - } - // Must write here, even if len(p) == 0, because it's "*one or more* - // s". - buf.WriteByte(byte(len(p))) - buf.Write(p) - return buf.Bytes() -} diff --git a/transport/internet/finalmask/xdns/dns_test.go b/transport/internet/finalmask/xdns/dns_test.go deleted file mode 100644 index 7eac084e6..000000000 --- a/transport/internet/finalmask/xdns/dns_test.go +++ /dev/null @@ -1,953 +0,0 @@ -package xdns - -import ( - "bytes" - "fmt" - "io" - "strconv" - "strings" - "testing" -) - -func namesEqual(a, b Name) bool { - if len(a) != len(b) { - return false - } - for i := 0; i < len(a); i++ { - if !bytes.Equal(a[i], b[i]) { - return false - } - } - return true -} - -func TestName(t *testing.T) { - for _, test := range []struct { - labels [][]byte - err error - s string - }{ - {[][]byte{}, nil, "."}, - {[][]byte{[]byte("test")}, nil, "test"}, - {[][]byte{[]byte("a"), []byte("b"), []byte("c")}, nil, "a.b.c"}, - - {[][]byte{{}}, ErrZeroLengthLabel, ""}, - {[][]byte{[]byte("a"), {}, []byte("c")}, ErrZeroLengthLabel, ""}, - - // 63 octets. - { - [][]byte{[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE")}, - nil, - "0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE", - }, - // 64 octets. - {[][]byte{[]byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDEF")}, ErrLabelTooLong, ""}, - - // 64+64+64+62 octets. - { - [][]byte{ - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABC"), - }, - nil, - "0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE.0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABC", - }, - // 64+64+64+63 octets. - {[][]byte{ - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCDE"), - []byte("0123456789abcdef0123456789ABCDEF0123456789abcdef0123456789ABCD"), - }, ErrNameTooLong, ""}, - // 127 one-octet labels. - { - [][]byte{ - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - }, - nil, - "0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E.F.0.1.2.3.4.5.6.7.8.9.a.b.c.d.e.f.0.1.2.3.4.5.6.7.8.9.A.B.C.D.E", - }, - // 128 one-octet labels. - {[][]byte{ - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'a'}, - {'b'}, - {'c'}, - {'d'}, - {'e'}, - {'f'}, - {'0'}, - {'1'}, - {'2'}, - {'3'}, - {'4'}, - {'5'}, - {'6'}, - {'7'}, - {'8'}, - {'9'}, - {'A'}, - {'B'}, - {'C'}, - {'D'}, - {'E'}, - {'F'}, - }, ErrNameTooLong, ""}, - } { - // Test that NewName returns proper error codes, and otherwise - // returns an equal slice of labels. - name, err := NewName(test.labels) - if err != test.err || (err == nil && !namesEqual(name, test.labels)) { - t.Errorf("%+q returned (%+q, %v), expected (%+q, %v)", - test.labels, name, err, test.labels, test.err) - continue - } - if test.err != nil { - continue - } - - // Test that the string version of the name comes out as - // expected. - s := name.String() - if s != test.s { - t.Errorf("%+q became string %+q, expected %+q", test.labels, s, test.s) - continue - } - - // Test that parsing from a string back to a Name results in the - // original slice of labels. - name, err = ParseName(s) - if err != nil || !namesEqual(name, test.labels) { - t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)", - test.labels, s, name, err, test.labels, nil) - continue - } - // A trailing dot should be ignored. - if !strings.HasSuffix(s, ".") { - dotName, dotErr := ParseName(s + ".") - if dotErr != err || !namesEqual(dotName, name) { - t.Errorf("%+q parsing %+q returned (%+q, %v), expected (%+q, %v)", - test.labels, s+".", dotName, dotErr, name, err) - continue - } - } - } -} - -func TestParseName(t *testing.T) { - for _, test := range []struct { - s string - name Name - err error - }{ - // This case can't be tested by TestName above because String - // will never produce "" (it produces "." instead). - {"", [][]byte{}, nil}, - } { - name, err := ParseName(test.s) - if err != test.err || (err == nil && !namesEqual(name, test.name)) { - t.Errorf("%+q returned (%+q, %v), expected (%+q, %v)", - test.s, name, err, test.name, test.err) - continue - } - } -} - -func unescapeString(s string) ([][]byte, error) { - if s == "." { - return [][]byte{}, nil - } - - var result [][]byte - for _, label := range strings.Split(s, ".") { - var buf bytes.Buffer - i := 0 - for i < len(label) { - switch label[i] { - case '\\': - if i+3 >= len(label) { - return nil, fmt.Errorf("truncated escape sequence at index %v", i) - } - if label[i+1] != 'x' { - return nil, fmt.Errorf("malformed escape sequence at index %v", i) - } - b, err := strconv.ParseUint(string(label[i+2:i+4]), 16, 8) - if err != nil { - return nil, fmt.Errorf("malformed hex sequence at index %v", i+2) - } - buf.WriteByte(byte(b)) - i += 4 - default: - buf.WriteByte(label[i]) - i++ - } - } - result = append(result, buf.Bytes()) - } - return result, nil -} - -func TestNameString(t *testing.T) { - for _, test := range []struct { - name Name - s string - }{ - {[][]byte{}, "."}, - {[][]byte{[]byte("\x00"), []byte("a.b"), []byte("c\nd\\")}, "\\x00.a\\x2eb.c\\x0ad\\x5c"}, - {[][]byte{ - []byte("\x00\x01\x02\x03\x04\x05\x06\x07\x08\t\n\x0b\x0c\r\x0e\x0f\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f !\"#$%&'()*+,-./0123456789:;<=>"), - []byte("?@ABCDEFGHIJKLMNOPQRSTUVWXYZ[\\]^_`abcdefghijklmnopqrstuvwxyz{|}"), - []byte("~\x7f\x80\x81\x82\x83\x84\x85\x86\x87\x88\x89\x8a\x8b\x8c\x8d\x8e\x8f\x90\x91\x92\x93\x94\x95\x96\x97\x98\x99\x9a\x9b\x9c\x9d\x9e\x9f\xa0\xa1\xa2\xa3\xa4\xa5\xa6\xa7\xa8\xa9\xaa\xab\xac\xad\xae\xaf\xb0\xb1\xb2\xb3\xb4\xb5\xb6\xb7\xb8\xb9\xba\xbb\xbc"), - []byte("\xbd\xbe\xbf\xc0\xc1\xc2\xc3\xc4\xc5\xc6\xc7\xc8\xc9\xca\xcb\xcc\xcd\xce\xcf\xd0\xd1\xd2\xd3\xd4\xd5\xd6\xd7\xd8\xd9\xda\xdb\xdc\xdd\xde\xdf\xe0\xe1\xe2\xe3\xe4\xe5\xe6\xe7\xe8\xe9\xea\xeb\xec\xed\xee\xef\xf0\xf1\xf2\xf3\xf4\xf5\xf6\xf7\xf8\xf9\xfa\xfb"), - []byte("\xfc\xfd\xfe\xff"), - }, "\\x00\\x01\\x02\\x03\\x04\\x05\\x06\\x07\\x08\\x09\\x0a\\x0b\\x0c\\x0d\\x0e\\x0f\\x10\\x11\\x12\\x13\\x14\\x15\\x16\\x17\\x18\\x19\\x1a\\x1b\\x1c\\x1d\\x1e\\x1f\\x20\\x21\\x22\\x23\\x24\\x25\\x26\\x27\\x28\\x29\\x2a\\x2b\\x2c-\\x2e\\x2f0123456789\\x3a\\x3b\\x3c\\x3d\\x3e.\\x3f\\x40ABCDEFGHIJKLMNOPQRSTUVWXYZ\\x5b\\x5c\\x5d\\x5e\\x5f\\x60abcdefghijklmnopqrstuvwxyz\\x7b\\x7c\\x7d.\\x7e\\x7f\\x80\\x81\\x82\\x83\\x84\\x85\\x86\\x87\\x88\\x89\\x8a\\x8b\\x8c\\x8d\\x8e\\x8f\\x90\\x91\\x92\\x93\\x94\\x95\\x96\\x97\\x98\\x99\\x9a\\x9b\\x9c\\x9d\\x9e\\x9f\\xa0\\xa1\\xa2\\xa3\\xa4\\xa5\\xa6\\xa7\\xa8\\xa9\\xaa\\xab\\xac\\xad\\xae\\xaf\\xb0\\xb1\\xb2\\xb3\\xb4\\xb5\\xb6\\xb7\\xb8\\xb9\\xba\\xbb\\xbc.\\xbd\\xbe\\xbf\\xc0\\xc1\\xc2\\xc3\\xc4\\xc5\\xc6\\xc7\\xc8\\xc9\\xca\\xcb\\xcc\\xcd\\xce\\xcf\\xd0\\xd1\\xd2\\xd3\\xd4\\xd5\\xd6\\xd7\\xd8\\xd9\\xda\\xdb\\xdc\\xdd\\xde\\xdf\\xe0\\xe1\\xe2\\xe3\\xe4\\xe5\\xe6\\xe7\\xe8\\xe9\\xea\\xeb\\xec\\xed\\xee\\xef\\xf0\\xf1\\xf2\\xf3\\xf4\\xf5\\xf6\\xf7\\xf8\\xf9\\xfa\\xfb.\\xfc\\xfd\\xfe\\xff"}, - } { - s := test.name.String() - if s != test.s { - t.Errorf("%+q escaped to %+q, expected %+q", test.name, s, test.s) - continue - } - unescaped, err := unescapeString(s) - if err != nil { - t.Errorf("%+q unescaping %+q resulted in error %v", test.name, s, err) - continue - } - if !namesEqual(Name(unescaped), test.name) { - t.Errorf("%+q roundtripped through %+q to %+q", test.name, s, unescaped) - continue - } - } -} - -func TestNameTrimSuffix(t *testing.T) { - for _, test := range []struct { - name, suffix string - trimmed string - ok bool - }{ - {"", "", ".", true}, - {".", ".", ".", true}, - {"abc", "", "abc", true}, - {"abc", ".", "abc", true}, - {"", "abc", ".", false}, - {".", "abc", ".", false}, - {"example.com", "com", "example", true}, - {"example.com", "net", ".", false}, - {"example.com", "example.com", ".", true}, - {"example.com", "test.com", ".", false}, - {"example.com", "xample.com", ".", false}, - {"example.com", "example", ".", false}, - {"example.com", "COM", "example", true}, - {"EXAMPLE.COM", "com", "EXAMPLE", true}, - } { - tmp, ok := mustParseName(test.name).TrimSuffix(mustParseName(test.suffix)) - trimmed := tmp.String() - if ok != test.ok || trimmed != test.trimmed { - t.Errorf("TrimSuffix %+q %+q returned (%+q, %v), expected (%+q, %v)", - test.name, test.suffix, trimmed, ok, test.trimmed, test.ok) - continue - } - } -} - -func TestReadName(t *testing.T) { - // Good tests. - for _, test := range []struct { - start int64 - end int64 - input string - s string - }{ - // Empty name. - {0, 1, "\x00abcd", "."}, - // No pointers. - {12, 25, "AAAABBBBCCCC\x07example\x03com\x00", "example.com"}, - // Backward pointer. - {25, 31, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x0c", "sub.example.com"}, - // Forward pointer. - {0, 4, "\x01a\xc0\x04\x03bcd\x00", "a.bcd"}, - // Two backwards pointers. - {31, 38, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x0c\x04sub2\xc0\x19", "sub2.sub.example.com"}, - // Forward then backward pointer. - {25, 31, "AAAABBBBCCCC\x07example\x03com\x00\x03sub\xc0\x1f\x04sub2\xc0\x0c", "sub.sub2.example.com"}, - // Overlapping codons. - {0, 4, "\x01a\xc0\x03bcd\x00", "a.bcd"}, - // Pointer to empty label. - {0, 10, "\x07example\xc0\x0a\x00", "example"}, - {1, 11, "\x00\x07example\xc0\x00", "example"}, - // Pointer to pointer to empty label. - {0, 10, "\x07example\xc0\x0a\xc0\x0c\x00", "example"}, - {1, 11, "\x00\x07example\xc0\x0c\xc0\x00", "example"}, - } { - r := bytes.NewReader([]byte(test.input)) - _, err := r.Seek(test.start, io.SeekStart) - if err != nil { - panic(err) - } - name, err := readName(r) - if err != nil { - t.Errorf("%+q returned error %s", test.input, err) - continue - } - s := name.String() - if s != test.s { - t.Errorf("%+q returned %+q, expected %+q", test.input, s, test.s) - continue - } - cur, _ := r.Seek(0, io.SeekCurrent) - if cur != test.end { - t.Errorf("%+q left offset %d, expected %d", test.input, cur, test.end) - continue - } - } - - // Bad tests. - for _, test := range []struct { - start int64 - input string - err error - }{ - {0, "", io.ErrUnexpectedEOF}, - // Reserved label type. - {0, "\x80example", ErrReservedLabelType}, - // Reserved label type. - {0, "\x40example", ErrReservedLabelType}, - // No Terminating empty label. - {0, "\x07example\x03com", io.ErrUnexpectedEOF}, - // Pointer past end of buffer. - {0, "\x07example\xc0\xff", io.ErrUnexpectedEOF}, - // Pointer to self. - {0, "\x07example\x03com\xc0\x0c", ErrTooManyPointers}, - // Pointer to self with intermediate label. - {0, "\x07example\x03com\xc0\x08", ErrTooManyPointers}, - // Two pointers that point to each other. - {0, "\xc0\x02\xc0\x00", ErrTooManyPointers}, - // Two pointers that point to each other, with intermediate labels. - {0, "\x01a\xc0\x04\x01b\xc0\x00", ErrTooManyPointers}, - // EOF while reading label. - {0, "\x0aexample", io.ErrUnexpectedEOF}, - // EOF before second byte of pointer. - {0, "\xc0", io.ErrUnexpectedEOF}, - {0, "\x07example\xc0", io.ErrUnexpectedEOF}, - } { - r := bytes.NewReader([]byte(test.input)) - _, err := r.Seek(test.start, io.SeekStart) - if err != nil { - panic(err) - } - name, err := readName(r) - if err == io.EOF { - err = io.ErrUnexpectedEOF - } - if err != test.err { - t.Errorf("%+q returned (%+q, %v), expected %v", test.input, name, err, test.err) - continue - } - } -} - -func mustParseName(s string) Name { - name, err := ParseName(s) - if err != nil { - panic(err) - } - return name -} - -func questionsEqual(a, b *Question) bool { - if !namesEqual(a.Name, b.Name) { - return false - } - if a.Type != b.Type || a.Class != b.Class { - return false - } - return true -} - -func rrsEqual(a, b *RR) bool { - if !namesEqual(a.Name, b.Name) { - return false - } - if a.Type != b.Type || a.Class != b.Class || a.TTL != b.TTL { - return false - } - if !bytes.Equal(a.Data, b.Data) { - return false - } - return true -} - -func messagesEqual(a, b *Message) bool { - if a.ID != b.ID || a.Flags != b.Flags { - return false - } - if len(a.Question) != len(b.Question) { - return false - } - for i := 0; i < len(a.Question); i++ { - if !questionsEqual(&a.Question[i], &b.Question[i]) { - return false - } - } - for _, rec := range []struct{ rrA, rrB []RR }{ - {a.Answer, b.Answer}, - {a.Authority, b.Authority}, - {a.Additional, b.Additional}, - } { - if len(rec.rrA) != len(rec.rrB) { - return false - } - for i := 0; i < len(rec.rrA); i++ { - if !rrsEqual(&rec.rrA[i], &rec.rrB[i]) { - return false - } - } - } - return true -} - -func TestMessageFromWireFormat(t *testing.T) { - for _, test := range []struct { - buf string - expected Message - err error - }{ - { - "\x12\x34", - Message{}, - io.ErrUnexpectedEOF, - }, - { - "\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01", - Message{ - ID: 0x1234, - Flags: 0x0100, - Question: []Question{ - { - Name: mustParseName("www.example.com"), - Type: 1, - Class: 1, - }, - }, - Answer: []RR{}, - Authority: []RR{}, - Additional: []RR{}, - }, - nil, - }, - { - "\x12\x34\x01\x00\x00\x01\x00\x00\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01X", - Message{}, - ErrTrailingBytes, - }, - { - "\x12\x34\x81\x80\x00\x01\x00\x01\x00\x00\x00\x00\x03www\x07example\x03com\x00\x00\x01\x00\x01\x03www\x07example\x03com\x00\x00\x01\x00\x01\x00\x00\x00\x80\x00\x04\xc0\x00\x02\x01", - Message{ - ID: 0x1234, - Flags: 0x8180, - Question: []Question{ - { - Name: mustParseName("www.example.com"), - Type: 1, - Class: 1, - }, - }, - Answer: []RR{ - { - Name: mustParseName("www.example.com"), - Type: 1, - Class: 1, - TTL: 128, - Data: []byte{192, 0, 2, 1}, - }, - }, - Authority: []RR{}, - Additional: []RR{}, - }, - nil, - }, - } { - message, err := MessageFromWireFormat([]byte(test.buf)) - if err != test.err || (err == nil && !messagesEqual(&message, &test.expected)) { - t.Errorf("%+q\nreturned (%+v, %v)\nexpected (%+v, %v)", - test.buf, message, err, test.expected, test.err) - continue - } - } -} - -func TestMessageWireFormatRoundTrip(t *testing.T) { - for _, message := range []Message{ - { - ID: 0x1234, - Flags: 0x0100, - Question: []Question{ - { - Name: mustParseName("www.example.com"), - Type: 1, - Class: 1, - }, - { - Name: mustParseName("www2.example.com"), - Type: 2, - Class: 2, - }, - }, - Answer: []RR{ - { - Name: mustParseName("abc"), - Type: 2, - Class: 3, - TTL: 0xffffffff, - Data: []byte{1}, - }, - { - Name: mustParseName("xyz"), - Type: 2, - Class: 3, - TTL: 255, - Data: []byte{}, - }, - }, - Authority: []RR{ - { - Name: mustParseName("."), - Type: 65535, - Class: 65535, - TTL: 0, - Data: []byte("XXXXXXXXXXXXXXXXXXX"), - }, - }, - Additional: []RR{}, - }, - } { - buf, err := message.WireFormat() - if err != nil { - t.Errorf("%+v cannot make wire format: %v", message, err) - continue - } - message2, err := MessageFromWireFormat(buf) - if err != nil { - t.Errorf("%+q cannot parse wire format: %v", buf, err) - continue - } - if !messagesEqual(&message, &message2) { - t.Errorf("messages unequal\nbefore: %+v\n after: %+v", message, message2) - continue - } - } -} - -func TestDecodeRDataTXT(t *testing.T) { - for _, test := range []struct { - p []byte - decoded []byte - err error - }{ - {[]byte{}, nil, io.ErrUnexpectedEOF}, - {[]byte("\x00"), []byte{}, nil}, - {[]byte("\x01"), nil, io.ErrUnexpectedEOF}, - } { - decoded, err := DecodeRDataTXT(test.p) - if err != test.err || (err == nil && !bytes.Equal(decoded, test.decoded)) { - t.Errorf("%+q\nreturned (%+q, %v)\nexpected (%+q, %v)", - test.p, decoded, err, test.decoded, test.err) - continue - } - } -} - -func TestEncodeRDataTXT(t *testing.T) { - // Encoding 0 bytes needs to return at least a single length octet of - // zero, not an empty slice. - p := make([]byte, 0) - encoded := EncodeRDataTXT(p) - if len(encoded) < 0 { - t.Errorf("EncodeRDataTXT(%v) returned %v", p, encoded) - } - - // 255 bytes should be able to be encoded into 256 bytes. - p = make([]byte, 255) - encoded = EncodeRDataTXT(p) - if len(encoded) > 256 { - t.Errorf("EncodeRDataTXT(%d bytes) returned %d bytes", len(p), len(encoded)) - } - - fmt.Println(EncodeRDataTXT(nil)) - fmt.Println(computeMaxEncodedPayload(maxUDPPayload)) -} - -func TestRDataTXTRoundTrip(t *testing.T) { - for _, p := range [][]byte{ - {}, - []byte("\x00"), - { - 0x00, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, 0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f, - 0x10, 0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18, 0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f, - 0x20, 0x21, 0x22, 0x23, 0x24, 0x25, 0x26, 0x27, 0x28, 0x29, 0x2a, 0x2b, 0x2c, 0x2d, 0x2e, 0x2f, - 0x30, 0x31, 0x32, 0x33, 0x34, 0x35, 0x36, 0x37, 0x38, 0x39, 0x3a, 0x3b, 0x3c, 0x3d, 0x3e, 0x3f, - 0x40, 0x41, 0x42, 0x43, 0x44, 0x45, 0x46, 0x47, 0x48, 0x49, 0x4a, 0x4b, 0x4c, 0x4d, 0x4e, 0x4f, - 0x50, 0x51, 0x52, 0x53, 0x54, 0x55, 0x56, 0x57, 0x58, 0x59, 0x5a, 0x5b, 0x5c, 0x5d, 0x5e, 0x5f, - 0x60, 0x61, 0x62, 0x63, 0x64, 0x65, 0x66, 0x67, 0x68, 0x69, 0x6a, 0x6b, 0x6c, 0x6d, 0x6e, 0x6f, - 0x70, 0x71, 0x72, 0x73, 0x74, 0x75, 0x76, 0x77, 0x78, 0x79, 0x7a, 0x7b, 0x7c, 0x7d, 0x7e, 0x7f, - 0x80, 0x81, 0x82, 0x83, 0x84, 0x85, 0x86, 0x87, 0x88, 0x89, 0x8a, 0x8b, 0x8c, 0x8d, 0x8e, 0x8f, - 0x90, 0x91, 0x92, 0x93, 0x94, 0x95, 0x96, 0x97, 0x98, 0x99, 0x9a, 0x9b, 0x9c, 0x9d, 0x9e, 0x9f, - 0xa0, 0xa1, 0xa2, 0xa3, 0xa4, 0xa5, 0xa6, 0xa7, 0xa8, 0xa9, 0xaa, 0xab, 0xac, 0xad, 0xae, 0xaf, - 0xb0, 0xb1, 0xb2, 0xb3, 0xb4, 0xb5, 0xb6, 0xb7, 0xb8, 0xb9, 0xba, 0xbb, 0xbc, 0xbd, 0xbe, 0xbf, - 0xc0, 0xc1, 0xc2, 0xc3, 0xc4, 0xc5, 0xc6, 0xc7, 0xc8, 0xc9, 0xca, 0xcb, 0xcc, 0xcd, 0xce, 0xcf, - 0xd0, 0xd1, 0xd2, 0xd3, 0xd4, 0xd5, 0xd6, 0xd7, 0xd8, 0xd9, 0xda, 0xdb, 0xdc, 0xdd, 0xde, 0xdf, - 0xe0, 0xe1, 0xe2, 0xe3, 0xe4, 0xe5, 0xe6, 0xe7, 0xe8, 0xe9, 0xea, 0xeb, 0xec, 0xed, 0xee, 0xef, - 0xf0, 0xf1, 0xf2, 0xf3, 0xf4, 0xf5, 0xf6, 0xf7, 0xf8, 0xf9, 0xfa, 0xfb, 0xfc, 0xfd, 0xfe, 0xff, - }, - } { - rdata := EncodeRDataTXT(p) - decoded, err := DecodeRDataTXT(rdata) - if err != nil || !bytes.Equal(decoded, p) { - t.Errorf("%+q returned (%+q, %v)", p, decoded, err) - continue - } - } -} - -func TestIPAnswerPayloadRoundTrip(t *testing.T) { - for _, rrType := range []uint16{RRTypeA, RRTypeAAAA} { - for _, payload := range [][]byte{ - {}, - {0x01}, - []byte("hello world"), - bytes.Repeat([]byte{0xab}, payloadChunkSizeForType(rrType)*3+1), - } { - question := Question{ - Name: mustParseName("example.com"), - Type: rrType, - Class: ClassIN, - } - answers, err := answersForPayload(question, responseTTL, payload) - if err != nil { - t.Fatalf("answersForPayload(%d) err = %v", rrType, err) - } - - if len(answers) > 1 { - answers[0], answers[len(answers)-1] = answers[len(answers)-1], answers[0] - } - - decoded := decodeResponsePayload(answers) - if !bytes.Equal(decoded, payload) { - t.Fatalf("rrType=%d decoded %x want %x", rrType, decoded, payload) - } - } - } -} - -func TestParseResolver(t *testing.T) { - tests := []struct { - resolver string - rrType uint16 - }{ - {"example.com+udp://1.1.1.1:53", RRTypeTXT}, - {"example.com:txt+udp://1.1.1.1:53", RRTypeTXT}, - {"example.com:a+udp://1.1.1.1:53", RRTypeA}, - {"example.com:aaaa+udp://1.1.1.1:53", RRTypeAAAA}, - } - - for _, test := range tests { - domain, server, rrType, err := parseResolver(test.resolver) - if err != nil { - t.Fatalf("parseResolver(%q) err = %v", test.resolver, err) - } - if domain.String() != "example.com" || server != "1.1.1.1:53" || rrType != test.rrType { - t.Fatalf("parseResolver(%q) = (%q, %q, %d)", test.resolver, domain.String(), server, rrType) - } - } -} - -func TestParseDomainSpec(t *testing.T) { - tests := []struct { - spec string - def string - rrType uint16 - wantErr bool - }{ - {"example.com", "", 0, false}, - {"example.com", "txt", RRTypeTXT, false}, - {"example.com:a", "", RRTypeA, false}, - {"example.com:aaaa", "", RRTypeAAAA, false}, - {"example.com:doh", "", 0, true}, - } - - for _, test := range tests { - got, err := parseDomainSpec(test.spec, test.def) - if test.wantErr { - if err == nil { - t.Fatalf("parseDomainSpec(%q, %q) err = nil", test.spec, test.def) - } - continue - } - if err != nil { - t.Fatalf("parseDomainSpec(%q, %q) err = %v", test.spec, test.def, err) - } - if got.name.String() != "example.com" || got.rrType != test.rrType { - t.Fatalf("parseDomainSpec(%q, %q) = (%q, %d)", test.spec, test.def, got.name.String(), got.rrType) - } - } -} - -func TestResponseForMethodRestriction(t *testing.T) { - query := &Message{ - ID: 1, - Flags: 0x0100, - Question: []Question{{ - Name: mustParseName("abc.example.com"), - Type: RRTypeTXT, - Class: ClassIN, - }}, - Additional: []RR{{ - Name: Name{}, - Type: RRTypeOPT, - Class: 4096, - }}, - } - - resp, _ := responseFor(query, []domainSpec{{name: mustParseName("example.com"), rrType: RRTypeA}}) - if resp == nil || resp.Rcode() != RcodeNameError { - t.Fatalf("responseFor method restriction rcode = %v", resp) - } - - resp, _ = responseFor(query, []domainSpec{{name: mustParseName("example.com")}}) - if resp == nil || resp.Rcode() != RcodeNoError { - t.Fatalf("responseFor unrestricted rcode = %v", resp) - } -} diff --git a/transport/internet/finalmask/xdns/domain.go b/transport/internet/finalmask/xdns/domain.go new file mode 100644 index 000000000..c7ca40870 --- /dev/null +++ b/transport/internet/finalmask/xdns/domain.go @@ -0,0 +1,215 @@ +package xdns + +import ( + "encoding/base32" + "errors" + "fmt" + "strings" + + "golang.org/x/net/dns/dnsmessage" + "golang.org/x/net/idna" +) + +func Lower(c byte) byte { + if c >= 'A' && c <= 'Z' { + return c + ('a' - 'A') + } + return c +} + +func ToUpper(b []byte) { + for i, c := range b { + if c >= 'a' && c <= 'z' { + b[i] = c - 'a' + 'A' + } + } +} + +func ToLower(b []byte) { + for i, c := range b { + if c >= 'A' && c <= 'Z' { + b[i] = c - 'A' + 'a' + } + } +} + +func NewTable() ([256]int, [256]int) { + var t, t_ [256]int + for i := range t { + t[i] = base32Encoding.DecodedLen(i) + } + for i := range t_ { + t_[i] = base32Encoding.EncodedLen(i) + } + return t, t_ +} + +const ( + TypeA uint16 = 1 + TypeCNAME uint16 = 5 + TypeTXT uint16 = 16 + TypeAAAA uint16 = 28 +) + +var ( + base32Encoding = base32.StdEncoding.WithPadding(base32.NoPadding) + table, table_ = NewTable() + TypeMap = map[uint16]byte{ + TypeA: 0, + TypeCNAME: 1, + TypeTXT: 2, + TypeAAAA: 3, + } + TypeMap_ = map[byte]uint16{ + 0: TypeA, + 1: TypeCNAME, + 2: TypeTXT, + 3: TypeAAAA, + } +) + +type Domain struct { + name dnsmessage.Name + lenLimit int + labelLimit int + types []uint16 + edns0 uint16 + + cap int + lenMax int +} + +func NewDomain(domain string, lenLimit int, labelLimit int, types []uint16, edns0 uint16) (*Domain, error) { + if strings.Contains(domain, "..") { + return nil, errors.New("invalid domain") + } + if lenLimit < 0 || lenLimit > 255 { + return nil, errors.New("lenLimit < 0 || lenLimit > 255") + } + if labelLimit < 0 || labelLimit > 63 { + return nil, errors.New("labelLimit < 0 || labelLimit > 63") + } + if len(types) == 0 { + return nil, errors.New("empty types") + } + for i := range types { + switch types[i] { + case uint16(dnsmessage.TypeA), uint16(dnsmessage.TypeCNAME), uint16(dnsmessage.TypeTXT), uint16(dnsmessage.TypeAAAA): + default: + return nil, errors.New("unknown types") + } + } + if edns0 != 0 && (edns0 < 512 || edns0 > 4096) { + return nil, errors.New("edns0 != 0 && (edns0 < 512 || edns0 > 4096)") + } + + ascii, err := idna.ToASCII(domain) + if err != nil { + return nil, err + } + ascii = strings.Trim(ascii, ".") + + name, err := dnsmessage.NewName(domain + ".") + if err != nil { + return nil, err + } + + if lenLimit < int(name.Length)+1 { + return nil, errors.New("lenLimit < int(name.Length)+1") + } + n := (lenLimit - int(name.Length) - 1) / (labelLimit + 1) + left := (lenLimit - int(name.Length) - 1) % (labelLimit + 1) + total := n * labelLimit + if left > 1 { + total += left - 1 + } + cap := table[total] + if cap < 17 { + return nil, errors.New("cap < 17") + } + total = table_[cap] + lenMax := int(name.Length) + 1 + total + total/labelLimit + if total%labelLimit > 0 { + lenMax += 1 + } + return &Domain{ + name: name, + lenLimit: lenLimit, + labelLimit: labelLimit, + types: types, + edns0: edns0, + + cap: cap, + lenMax: lenMax, + }, nil +} + +func (d *Domain) Show() string { + return fmt.Sprint(d.name, d.cap) +} + +func (d *Domain) IsDomain(name dnsmessage.Name) bool { + if d.name.Length >= name.Length { + return false + } + i := d.name.Length + j := name.Length + for i > 0 { + i-- + j-- + if Lower(d.name.Data[i]) != Lower(name.Data[j]) { + return false + } + } + return true +} + +func (d *Domain) HasType(qtype uint16) bool { + for i := range d.types { + if d.types[i] == qtype { + return true + } + } + return false +} + +func (d *Domain) Encode(data []byte) dnsmessage.Name { + var name dnsmessage.Name + var encoded [255]byte + base32Encoding.Encode(encoded[:], data) + ToLower(encoded[:table_[len(data)]]) + b1 := name.Data[:0] + b2 := encoded[:table_[len(data)]] + for len(b2) > 0 { + size := min(len(b2), d.labelLimit) + b1 = append(b1, b2[:size]...) + b1 = append(b1, '.') + b2 = b2[size:] + } + b1 = append(b1, d.name.Data[:d.name.Length]...) + if len(b1) > 254 { + panic("len(b1) > 254") + } + name.Length = byte(len(b1)) + return name +} + +func (d *Domain) Decode(decoded *[255]byte, name dnsmessage.Name) int { + if !d.IsDomain(name) { + return 0 + } + var encoded [255]byte + b1 := encoded[:0] + b2 := name.Data[:name.Length-d.name.Length] + for i := range b2 { + if b2[i] != '.' { + b1 = append(b1, b2[i]) + } + } + ToUpper(b1) + n, err := base32Encoding.Decode(decoded[:], b1) + if err != nil { + return 0 + } + return n +} diff --git a/transport/internet/finalmask/xdns/frag.go b/transport/internet/finalmask/xdns/frag.go new file mode 100644 index 000000000..444f405af --- /dev/null +++ b/transport/internet/finalmask/xdns/frag.go @@ -0,0 +1,171 @@ +package xdns + +import ( + "sync" + "time" +) + +const ( + fragTTL = 8 * time.Second + fragSize = 4096 + fragClientIDSize = 16384 + fragCount = 4096 +) + +type FragKey struct { + clientID ClientID + fragID byte +} + +type FragEntry struct { + data [][]byte + size int + len int + total byte + deadline time.Time +} + +type FragManager struct { + m map[FragKey]*FragEntry + sizem map[ClientID]int + ch chan struct{} + mu sync.Mutex +} + +func NewFragManager() *FragManager { + m := &FragManager{ + m: make(map[FragKey]*FragEntry), + sizem: make(map[ClientID]int), + ch: make(chan struct{}), + } + go m.gc() + return m +} + +func (m *FragManager) closed() bool { + select { + case <-m.ch: + return true + default: + return false + } +} + +func (m *FragManager) removeEntey(k FragKey, e *FragEntry) { + m.sizem[k.clientID] -= e.size + delete(m.m, k) +} + +func (m *FragManager) tryRemove() { + if len(m.m) < fragCount { + return + } + var key FragKey + var entry *FragEntry + first := true + for k, e := range m.m { + if first || e.deadline.Before(entry.deadline) { + key = k + entry = e + first = false + } + } + m.removeEntey(key, entry) +} + +func (m *FragManager) gc() { + ticker := time.NewTicker(fragTTL / 2) + defer ticker.Stop() + for { + select { + case <-m.ch: + return + case now := <-ticker.C: + m.mu.Lock() + for k, e := range m.m { + if now.After(e.deadline) { + m.removeEntey(k, e) + } + } + m.mu.Unlock() + } + } +} + +func (m *FragManager) Feed(out []byte, key FragKey, fragIdx, fragN byte, data []byte) int { + m.mu.Lock() + defer m.mu.Unlock() + if m.closed() { + return 0 + } + + if fragN < 2 { + return 0 + } + + now := time.Now() + entry := m.m[key] + if entry == nil || now.After(entry.deadline) { + if entry == nil { + m.tryRemove() + } else { + m.removeEntey(key, entry) + } + entry = &FragEntry{ + data: make([][]byte, fragN), + total: fragN, + deadline: now.Add(fragTTL), + } + m.m[key] = entry + } + + if fragN != entry.total { + return 0 + } + if fragIdx >= entry.total { + return 0 + } + if entry.data[fragIdx] != nil { + return 0 + } + if entry.size+len(data) > fragSize { + return 0 + } + if entry.len < int(entry.total)-1 { + if m.sizem[key.clientID]+len(data) > fragClientIDSize { + return 0 + } + } + + cp := make([]byte, len(data)) + copy(cp, data) + + entry.data[fragIdx] = cp + entry.size += len(data) + entry.len++ + entry.deadline = now.Add(fragTTL) + m.sizem[key.clientID] += len(data) + + if entry.len < int(entry.total) { + return 0 + } + + out = out[:0] + for i := range entry.data { + out = append(out, entry.data[i]...) + } + m.removeEntey(key, entry) + return len(out) +} + +func (m *FragManager) Close() { + m.mu.Lock() + defer m.mu.Unlock() + if m.closed() { + return + } + close(m.ch) + for k := range m.m { + delete(m.m, k) + } +} diff --git a/transport/internet/finalmask/xdns/record_transport.go b/transport/internet/finalmask/xdns/record_transport.go deleted file mode 100644 index 8428baa40..000000000 --- a/transport/internet/finalmask/xdns/record_transport.go +++ /dev/null @@ -1,226 +0,0 @@ -package xdns - -import "bytes" - -const ipRecordHeaderSize = 2 - -func maxEncodedPayloadForType(rrType uint16) int { - switch rrType { - case RRTypeA: - return maxEncodedPayloadA - case RRTypeAAAA: - return maxEncodedPayloadAAAA - default: - return maxEncodedPayloadTXT - } -} - -func rrDataSizeForType(rrType uint16) int { - switch rrType { - case RRTypeA: - return 4 - case RRTypeAAAA: - return 16 - default: - return 0 - } -} - -func payloadChunkSizeForType(rrType uint16) int { - size := rrDataSizeForType(rrType) - if size <= ipRecordHeaderSize { - return 0 - } - return size - ipRecordHeaderSize -} - -func answersForPayload(question Question, ttl uint32, payload []byte) ([]RR, error) { - switch question.Type { - case RRTypeTXT: - return []RR{ - { - Name: question.Name, - Type: question.Type, - Class: question.Class, - TTL: ttl, - Data: EncodeRDataTXT(payload), - }, - }, nil - case RRTypeA, RRTypeAAAA: - return ipAnswersForPayload(question, ttl, payload) - default: - return nil, ErrIntegerOverflow - } -} - -func ipAnswersForPayload(question Question, ttl uint32, payload []byte) ([]RR, error) { - chunkSize := payloadChunkSizeForType(question.Type) - rrDataSize := rrDataSizeForType(question.Type) - if chunkSize == 0 || rrDataSize == 0 { - return nil, ErrIntegerOverflow - } - - numRecords := 1 - if len(payload) > 0 { - numRecords = (len(payload) + chunkSize - 1) / chunkSize - } - if numRecords > 256 { - return nil, ErrIntegerOverflow - } - - answers := make([]RR, 0, numRecords) - for i := 0; i < numRecords; i++ { - offset := i * chunkSize - n := len(payload) - offset - if n < 0 { - n = 0 - } - if n > chunkSize { - n = chunkSize - } - - data := make([]byte, rrDataSize) - data[0] = byte(i) - data[1] = byte(n) - copy(data[ipRecordHeaderSize:], payload[offset:offset+n]) - - answers = append(answers, RR{ - Name: question.Name, - Type: question.Type, - Class: question.Class, - TTL: ttl, - Data: data, - }) - } - - return answers, nil -} - -func decodeResponsePayload(answers []RR) []byte { - if len(answers) == 0 { - return nil - } - - switch answers[0].Type { - case RRTypeTXT: - if len(answers) != 1 { - return nil - } - payload, err := DecodeRDataTXT(answers[0].Data) - if err != nil { - return nil - } - return payload - case RRTypeA, RRTypeAAAA: - return decodeIPAnswerPayload(answers, answers[0].Type) - default: - return nil - } -} - -func decodeIPAnswerPayload(answers []RR, rrType uint16) []byte { - chunkSize := payloadChunkSizeForType(rrType) - rrDataSize := rrDataSizeForType(rrType) - if chunkSize == 0 || rrDataSize == 0 || len(answers) > 256 { - return nil - } - - parts := make([][]byte, len(answers)) - for _, answer := range answers { - if answer.Type != rrType || len(answer.Data) != rrDataSize { - return nil - } - idx := int(answer.Data[0]) - n := int(answer.Data[1]) - if idx >= len(answers) || n > chunkSize || parts[idx] != nil { - return nil - } - - part := make([]byte, n) - copy(part, answer.Data[ipRecordHeaderSize:ipRecordHeaderSize+n]) - parts[idx] = part - } - - var payload bytes.Buffer - for _, part := range parts { - if part == nil { - return nil - } - payload.Write(part) - } - return payload.Bytes() -} - -func computeMaxEncodedPayload(limit int) int { - return computeMaxEncodedPayloadForType(limit, RRTypeTXT) -} - -func computeMaxEncodedPayloadForType(limit int, rrType uint16) int { - maxLengthName, err := NewName([][]byte{ - []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), - []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), - []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), - []byte("AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"), - }) - if err != nil { - panic(err) - } - { - n := 0 - for _, label := range maxLengthName { - n += len(label) + 1 - } - n += 1 - if n != 255 { - panic("computeMaxEncodedPayload n != 255") - } - } - - queryLimit := uint16(limit) - if int(queryLimit) != limit { - queryLimit = 0xffff - } - query := &Message{ - Question: []Question{ - { - Name: maxLengthName, - Type: rrType, - Class: ClassIN, - }, - }, - Additional: []RR{ - { - Name: Name{}, - Type: RRTypeOPT, - Class: queryLimit, - TTL: 0, - Data: []byte{}, - }, - }, - } - resp, _ := responseFor(query, []domainSpec{{name: Name{[]byte{}}}}) - - low := 0 - high := 32768 - if chunkSize := payloadChunkSizeForType(rrType); chunkSize > 0 { - high = 256*chunkSize + 1 - } - for low+1 < high { - mid := (low + high) / 2 - resp.Answer, err = answersForPayload(query.Question[0], responseTTL, make([]byte, mid)) - if err != nil { - panic(err) - } - buf, err := resp.WireFormat() - if err != nil { - panic(err) - } - if len(buf) <= limit { - low = mid - } else { - high = mid - } - } - - return low -} diff --git a/transport/internet/finalmask/xdns/resolver.go b/transport/internet/finalmask/xdns/resolver.go new file mode 100644 index 000000000..279af5053 --- /dev/null +++ b/transport/internet/finalmask/xdns/resolver.go @@ -0,0 +1,31 @@ +package xdns + +import ( + "errors" + "net" + + "github.com/xtls/xray-core/common/serial" + "github.com/xtls/xray-core/transport/internet/finalmask" +) + +type Resolver interface { + Addr() *net.UDPAddr + Read(p []byte) (int, error) + Send(p []byte) + Close() +} + +func NewResolver(proto *serial.TypedMessage, dialer *finalmask.Dialer) (Resolver, error) { + config, err := proto.GetInstance() + if err != nil { + return nil, err + } + switch v := config.(type) { + case *TCPResolverProto: + return NewTCPResolver(v, dialer) + case *UDPResolverProto: + return NewUDPResolver(v, dialer) + default: + return nil, errors.New("unknown proto") + } +} diff --git a/transport/internet/finalmask/xdns/resolver_tcp.go b/transport/internet/finalmask/xdns/resolver_tcp.go new file mode 100644 index 000000000..89f2b0ea6 --- /dev/null +++ b/transport/internet/finalmask/xdns/resolver_tcp.go @@ -0,0 +1,143 @@ +package xdns + +import ( + "encoding/binary" + "errors" + "io" + "sync" + + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" +) + +type TCPResolver struct { + dest net.Destination + dialer *finalmask.Dialer + + conn net.Conn + tcpAddr *net.TCPAddr + udpAddr *net.UDPAddr + + readCh chan []byte + closeCh chan struct{} + wg sync.WaitGroup + mu sync.Mutex +} + +func NewTCPResolver(config *TCPResolverProto, dialer *finalmask.Dialer) (Resolver, error) { + dest, err := net.ParseDestination("tcp:" + config.Addr) + if err != nil { + return nil, err + } + r := &TCPResolver{ + dest: dest, + dialer: dialer, + readCh: make(chan []byte), + closeCh: make(chan struct{}), + } + if err := r.dial(); err != nil { + r.Close() + return nil, err + } + return r, nil +} + +func (r *TCPResolver) closed() bool { + select { + case <-r.closeCh: + return true + default: + return false + } +} + +func (r *TCPResolver) dial() error { + if r.closed() { + return errors.New("closed") + } + if r.conn != nil { + return nil + } + conn, err := r.dialer.DialTCP(r.dest) + if err != nil { + return err + } + r.conn = conn + r.tcpAddr = conn.RemoteAddr().(*net.TCPAddr) + r.udpAddr = &net.UDPAddr{IP: r.tcpAddr.IP, Port: r.tcpAddr.Port} + r.wg.Add(1) + go r.recv(conn) + return nil +} + +func (r *TCPResolver) recv(conn net.Conn) { + defer r.wg.Done() + + var buf [4096]byte + for { + _, err := io.ReadFull(conn, buf[:2]) + if err != nil { + break + } + n := binary.BigEndian.Uint16(buf[:2]) + if n == 0 || n > 4096 { + io.CopyN(io.Discard, conn, int64(n)) + continue + } + _, err = io.ReadFull(conn, buf[:n]) + if err != nil { + break + } + p := pool4K.Get().([]byte) + copy(p, buf[:n]) + select { + case <-r.closeCh: + pool4K.Put(p[:cap(p)]) + case r.readCh <- p[:n]: + } + } + + r.mu.Lock() + defer r.mu.Unlock() + + _ = conn.Close() + r.conn = nil +} + +func (r *TCPResolver) Addr() *net.UDPAddr { + return r.udpAddr +} + +func (r *TCPResolver) Read(p []byte) (n int, err error) { + packet, ok := <-r.readCh + if ok { + n = copy(p, packet) + pool4K.Put(packet[:cap(packet)]) + return n, nil + } + return 0, io.ErrClosedPipe +} + +func (r *TCPResolver) Send(p []byte) { + r.mu.Lock() + defer r.mu.Unlock() + if r.dial() != nil { + return + } + _ = binary.Write(r.conn, binary.BigEndian, len(p)) + _, _ = r.conn.Write(p) +} + +func (r *TCPResolver) Close() { + r.mu.Lock() + defer r.mu.Unlock() + if r.closed() { + return + } + close(r.closeCh) + if r.conn != nil { + _ = r.conn.Close() + } + r.wg.Wait() + close(r.readCh) +} diff --git a/transport/internet/finalmask/xdns/resolver_udp.go b/transport/internet/finalmask/xdns/resolver_udp.go new file mode 100644 index 000000000..e8d97ba29 --- /dev/null +++ b/transport/internet/finalmask/xdns/resolver_udp.go @@ -0,0 +1,130 @@ +package xdns + +import ( + "errors" + "io" + "sync" + + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet/finalmask" +) + +type UDPResolver struct { + dest net.Destination + dialer *finalmask.Dialer + + conn net.PacketConn + udpAddr *net.UDPAddr + + readCh chan []byte + closeCh chan struct{} + wg sync.WaitGroup + mu sync.Mutex +} + +func NewUDPResolver(config *UDPResolverProto, dialer *finalmask.Dialer) (Resolver, error) { + dest, err := net.ParseDestination("udp:" + config.Addr) + if err != nil { + return nil, err + } + r := &UDPResolver{ + dest: dest, + dialer: dialer, + readCh: make(chan []byte), + closeCh: make(chan struct{}), + } + if err := r.dial(); err != nil { + r.Close() + return nil, err + } + return r, nil +} + +func (r *UDPResolver) closed() bool { + select { + case <-r.closeCh: + return true + default: + return false + } +} + +func (r *UDPResolver) dial() error { + if r.closed() { + return errors.New("closed") + } + if r.conn != nil { + return nil + } + conn, err := r.dialer.DialUDP(r.dest) + if err != nil { + return err + } + r.conn = conn.(*net.PacketConnWrapper).PacketConn + r.udpAddr = conn.RemoteAddr().(*net.UDPAddr) + r.wg.Add(1) + go r.recv(conn.(*net.PacketConnWrapper).PacketConn) + return nil +} + +func (r *UDPResolver) recv(conn net.PacketConn) { + defer r.wg.Done() + + var buf [4096]byte + for { + n, _, err := conn.ReadFrom(buf[:]) + if err != nil { + break + } + p := pool4K.Get().([]byte) + copy(p, buf[:n]) + select { + case <-r.closeCh: + pool4K.Put(p[:cap(p)]) + case r.readCh <- p[:n]: + } + } + + r.mu.Lock() + defer r.mu.Unlock() + + _ = conn.Close() + r.conn = nil +} + +func (r *UDPResolver) Addr() *net.UDPAddr { + return r.udpAddr +} + +func (r *UDPResolver) Read(p []byte) (n int, err error) { + packet, ok := <-r.readCh + if ok { + n = copy(p, packet) + pool4K.Put(packet[:cap(packet)]) + return n, nil + } + return 0, io.ErrClosedPipe +} + +func (r *UDPResolver) Send(p []byte) { + r.mu.Lock() + defer r.mu.Unlock() + if err := r.dial(); err != nil { + return + } + _, _ = r.conn.WriteTo(p, r.udpAddr) +} + +func (r *UDPResolver) Close() { + r.mu.Lock() + defer r.mu.Unlock() + if r.closed() { + return + } + close(r.closeCh) + if r.conn != nil { + _ = r.conn.Close() + } + r.wg.Wait() + close(r.readCh) +} diff --git a/transport/internet/finalmask/xdns/resp.go b/transport/internet/finalmask/xdns/resp.go new file mode 100644 index 000000000..93ff8a8c3 --- /dev/null +++ b/transport/internet/finalmask/xdns/resp.go @@ -0,0 +1,392 @@ +package xdns + +import ( + "sort" + "sync" + "time" + + "github.com/xtls/xray-core/common" + "golang.org/x/net/dns/dnsmessage" +) + +const ( + sendTTL = 4 * time.Second +) + +type Resp struct { + msg dnsmessage.Message + domain *Domain + edns0 uint16 + + cap int +} + +func NewResp(msg dnsmessage.Message, domain *Domain, edns0 uint16) *Resp { + if msg.Header.Response { + return &Resp{ + msg: msg, + domain: domain, + } + } + + size := min(max(int(edns0), 512), max(int(domain.edns0), 512)) + + left := size - 12 - int(msg.Questions[0].Name.Length) - 1 - 2 - 2 + if edns0 > 0 { + left -= 1 + 2 + 2 + 4 + 2 + 0 + } + cap := 0 + switch msg.Questions[0].Type { + case dnsmessage.TypeA: + single := 2 + 2 + 2 + 4 + 2 + 4 + n := left / single + if n > 255 { + n = 255 + } + cap = 4*n - n - 1 + case dnsmessage.TypeCNAME: + single := 2 + 2 + 2 + 4 + 2 + domain.lenMax + n := left / single + if n > 255 { + n = 255 + } + cap = domain.cap*n - n - 1 + case dnsmessage.TypeTXT: + left -= 2 + 2 + 2 + 4 + 2 + single := 255 + n := left / single + m := left % single + cap = 255*n - n + if m > 1 { + cap += m - 1 + } + case dnsmessage.TypeAAAA: + single := 2 + 2 + 2 + 4 + 2 + 16 + n := left / single + if n > 255 { + n = 255 + } + cap = 16*n - n - 1 + } + + return &Resp{ + msg: msg, + domain: domain, + edns0: edns0, + + cap: cap, + } +} + +func (r *Resp) Encode(encoded []byte, data []byte) []byte { + msg := r.msg + msg.Header = dnsmessage.Header{ + ID: msg.Header.ID, + Response: true, + Authoritative: true, + RCode: dnsmessage.RCodeSuccess, + } + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + switch msg.Questions[0].Type { + case dnsmessage.TypeA: + fragN := 0 + if len(data) > 0 { + fragN = 1 + } + if (len(data) - (4 - 2)) > 0 { + fragN += (len(data) - (4 - 2)) / (4 - 1) + if (len(data)-(4-2))%(4-1) > 0 { + fragN++ + } + } + + for i := range fragN { + A := [4]byte{byte(i)} + if i == 0 { + A[1] = byte(fragN) + n := copy(A[2:], data) + data = data[n:] + } else { + n := copy(A[1:], data) + data = data[n:] + } + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.AResource{A: A}, + }) + } + case dnsmessage.TypeCNAME: + fragN := 0 + if len(data) > 0 { + fragN = 1 + } + if (len(data) - (r.domain.cap - 2)) > 0 { + fragN += (len(data) - (r.domain.cap - 2)) / (r.domain.cap - 1) + if (len(data)-(r.domain.cap-2))%(r.domain.cap-1) > 0 { + fragN++ + } + } + + DATA := make([]byte, r.domain.cap) + for i := range fragN { + DATA[0] = byte(i) + if i == 0 { + DATA[1] = byte(fragN) + n := copy(DATA[2:], data) + data = data[n:] + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.CNAMEResource{CNAME: r.domain.Encode(DATA[:2+n])}, + }) + } else { + n := copy(DATA[1:], data) + data = data[n:] + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.CNAMEResource{CNAME: r.domain.Encode(DATA[:1+n])}, + }) + } + } + case dnsmessage.TypeTXT: + var txt []string + for len(data) > 0 { + size := min(len(data), 255) + txt = append(txt, string(data[:size])) + data = data[size:] + } + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.TXTResource{TXT: txt}, + }) + case dnsmessage.TypeAAAA: + fragN := 0 + if len(data) > 0 { + fragN = 1 + } + if (len(data) - (16 - 2)) > 0 { + fragN += (len(data) - (16 - 2)) / (16 - 1) + if (len(data)-(16-2))%(16-1) > 0 { + fragN++ + } + } + + for i := range fragN { + AAAA := [16]byte{byte(i)} + if i == 0 { + AAAA[1] = byte(fragN) + n := copy(AAAA[2:], data) + data = data[n:] + } else { + n := copy(AAAA[1:], data) + data = data[n:] + } + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: msg.Questions[0].Name, + Type: msg.Questions[0].Type, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.AAAAResource{AAAA: AAAA}, + }) + } + } + if r.edns0 > 0 { + msg.Additionals = append(msg.Additionals, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("."), + Type: dnsmessage.TypeOPT, + Class: dnsmessage.Class(r.edns0), + TTL: 0, + }, + Body: &dnsmessage.OPTResource{}, + }) + } + return common.Must2(msg.AppendPack(encoded[:0])) +} + +func (r *Resp) Decode(decoded []byte) int { + decoded = decoded[:0] + msg := r.msg + if msg.Questions[0].Type == dnsmessage.TypeTXT { + if len(msg.Answers) == 1 && r.domain.IsDomain(msg.Answers[0].Header.Name) && msg.Answers[0].Header.Type == dnsmessage.TypeTXT { + for i := range msg.Answers[0].Body.(*dnsmessage.TXTResource).TXT { + decoded = append(decoded, msg.Answers[0].Body.(*dnsmessage.TXTResource).TXT[i]...) + } + } + return len(decoded) + } else { + var frags [][]byte + for i := range msg.Answers { + if !r.domain.IsDomain(msg.Answers[i].Header.Name) || msg.Answers[i].Header.Type != msg.Questions[0].Type { + continue + } + switch msg.Questions[0].Type { + case dnsmessage.TypeA: + frags = append(frags, msg.Answers[i].Body.(*dnsmessage.AResource).A[:]) + case dnsmessage.TypeCNAME: + var decoded [255]byte + n := r.domain.Decode(&decoded, msg.Answers[i].Body.(*dnsmessage.CNAMEResource).CNAME) + if n == 0 { + continue + } + frags = append(frags, decoded[:n]) + case dnsmessage.TypeAAAA: + frags = append(frags, msg.Answers[i].Body.(*dnsmessage.AAAAResource).AAAA[:]) + } + } + sort.Slice(frags, func(i, j int) bool { + return frags[i][0] < frags[j][0] + }) + if len(frags) < 1 || len(frags[0]) < 2 || int(frags[0][1]) > len(frags) { + return 0 + } + decoded = append(decoded, frags[0][2:]...) + for i := range frags { + if i > 0 { + if frags[i][0] == frags[i-1][0] { + return 0 + } + decoded = append(decoded, frags[i][1:]...) + } + } + return len(decoded) + } +} + +type SendInfo struct { + stash chan []byte + ch chan []byte + deadline time.Time +} + +type SendManager struct { + m map[ClientID]*SendInfo + ch chan struct{} + mu sync.Mutex +} + +func NewSendManager() *SendManager { + m := &SendManager{ + m: make(map[ClientID]*SendInfo), + ch: make(chan struct{}), + } + go m.gc() + return m +} + +func (m *SendManager) closed() bool { + select { + case <-m.ch: + return true + default: + return false + } +} + +func (m *SendManager) gc() { + ticker := time.NewTicker(sendTTL) + defer ticker.Stop() + for { + select { + case <-m.ch: + return + case now := <-ticker.C: + m.mu.Lock() + for key, info := range m.m { + if now.After(info.deadline) { + close(info.stash) + close(info.ch) + delete(m.m, key) + } + } + m.mu.Unlock() + ticker.Reset(sendTTL) + } + } +} + +func (m *SendManager) Push(clientID ClientID, p []byte) { + m.mu.Lock() + defer m.mu.Unlock() + info := m.m[clientID] + if info == nil { + info = &SendInfo{ + stash: make(chan []byte, 1), + ch: make(chan []byte, 128), + deadline: time.Now().Add(sendTTL), + } + m.m[clientID] = info + } + b := make([]byte, len(p)) + copy(b, p) + select { + case info.ch <- b: + default: + } +} + +func (m *SendManager) Stash(clientID ClientID, p []byte) { + m.mu.Lock() + defer m.mu.Unlock() + info := m.m[clientID] + if info == nil { + return + } + info.deadline = time.Now().Add(sendTTL) + select { + case info.stash <- p: + default: + } +} + +func (m *SendManager) Pop(clientID ClientID) (chan []byte, chan []byte) { + m.mu.Lock() + defer m.mu.Unlock() + info := m.m[clientID] + if info == nil { + info = &SendInfo{ + stash: make(chan []byte, 1), + ch: make(chan []byte, 128), + } + m.m[clientID] = info + } + info.deadline = time.Now().Add(sendTTL) + return info.ch, info.stash +} + +func (m *SendManager) Close() { + m.mu.Lock() + defer m.mu.Unlock() + if m.closed() { + return + } + close(m.ch) + for key, info := range m.m { + close(info.stash) + close(info.ch) + delete(m.m, key) + } +} diff --git a/transport/internet/finalmask/xdns/server.go b/transport/internet/finalmask/xdns/server.go index 654f7fdba..c7506f5f3 100644 --- a/transport/internet/finalmask/xdns/server.go +++ b/transport/internet/finalmask/xdns/server.go @@ -1,512 +1,385 @@ package xdns import ( - "bytes" "context" - "encoding/binary" - go_errors "errors" "io" - "net" "sync" "time" + "github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common/errors" - "github.com/xtls/xray-core/transport/internet/finalmask" + "github.com/xtls/xray-core/common/net" + "golang.org/x/net/dns/dnsmessage" ) const ( - idleTimeout = 10 * time.Second - responseTTL = 60 - maxResponseDelay = 1 * time.Second + maxResponseDelay = time.Second ) -var ( - maxUDPPayload = 1280 - 40 - 8 - maxEncodedPayloadTXT = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeTXT) - maxEncodedPayloadA = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeA) - maxEncodedPayloadAAAA = computeMaxEncodedPayloadForType(maxUDPPayload, RRTypeAAAA) -) - -func clientIDToAddr(clientID [8]byte) *net.UDPAddr { - ip := make(net.IP, 16) - - copy(ip, []byte{0xfd, 0x00, 0, 0, 0, 0, 0, 0}) - copy(ip[8:], clientID[:]) - - return &net.UDPAddr{ - IP: ip, - } +type resp struct { + msg dnsmessage.Message + addr net.Addr } -type record struct { - Resp *Message - Addr net.Addr - // ClientID [8]byte - ClientAddr net.Addr +type Rec struct { + resp *Resp + clientID ClientID + addr net.Addr } -type queue struct { - last time.Time - rrType uint16 - queue chan []byte - stash chan []byte -} - -type xdnsConnServer struct { +type xdnsServer struct { net.PacketConn - domains []domainSpec + domains []*Domain + fragManager *FragManager + sendManager *SendManager - ch chan *record - readQueue chan *packet - writeQueueMap map[string]*queue - - closed bool - mutex sync.Mutex + readCh chan packet + recCh chan *Rec + drCh chan resp + closeCh chan struct{} + wg sync.WaitGroup + mu sync.RWMutex } -func NewConnServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { +func NewServer(c *Config, raw net.PacketConn) (net.PacketConn, error) { if len(c.Domains) == 0 { return nil, errors.New("empty domains") } - domains := make([]domainSpec, 0, len(c.Domains)) - for _, domain := range c.Domains { - domain, err := parseDomainSpec(domain, "") + domains := make([]*Domain, 0, len(c.Domains)) + for i := range c.Domains { + 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 := 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 } domains = append(domains, domain) } - - conn := &xdnsConnServer{ + server := &xdnsServer{ PacketConn: raw, - domains: domains, + domains: domains, + fragManager: NewFragManager(), + sendManager: NewSendManager(), - ch: make(chan *record, 500), - readQueue: make(chan *packet, 512), - writeQueueMap: make(map[string]*queue), + readCh: make(chan packet), + recCh: make(chan *Rec, 255), + drCh: make(chan resp), + closeCh: make(chan struct{}), } - - go conn.clean() - go conn.recvLoop() - go conn.sendLoop() - - return conn, nil + go server.run() + return server, nil } -func (c *xdnsConnServer) clean() { - f := func() bool { - c.mutex.Lock() - defer c.mutex.Unlock() - - if c.closed { - return true - } - - now := time.Now() - - for key, q := range c.writeQueueMap { - if now.Sub(q.last) >= idleTimeout { - close(q.queue) - close(q.stash) - delete(c.writeQueueMap, key) - } - } - +func (c *xdnsServer) closed() bool { + select { + case <-c.closeCh: + return true + default: return false } - - for { - time.Sleep(idleTimeout / 2) - if f() { - return - } - } } -func (c *xdnsConnServer) ensureQueue(addr net.Addr) *queue { - if c.closed { - return nil - } - - q, ok := c.writeQueueMap[addr.String()] - if !ok { - q = &queue{ - queue: make(chan []byte, 512), - stash: make(chan []byte, 1), - } - c.writeQueueMap[addr.String()] = q - } - q.last = time.Now() - - return q -} - -func (c *xdnsConnServer) stash(queue *queue, p []byte) { - c.mutex.Lock() - defer c.mutex.Unlock() - - if c.closed { - return - } - +func (c *xdnsServer) decref(msg dnsmessage.Message, addr net.Addr) { select { - case queue.stash <- p: + case c.drCh <- resp{msg: msg, addr: addr}: default: } } -func (c *xdnsConnServer) recvLoop() { - var buf [finalmask.UDPSize]byte +func (c *xdnsServer) read(buf []byte, addr net.Addr) { + msg := dnsmessage.Message{} + if err := msg.Unpack(buf); err != nil { + return + } + if msg.Header.Response { + return + } - for { - if c.closed { - break - } + if msg.Header.OpCode != 0 { + msg.Header.Response = true + msg.Header.RCode = dnsmessage.RCodeNotImplemented + c.decref(msg, addr) + return + } - n, addr, err := c.PacketConn.ReadFrom(buf[:]) - if err != nil { - if go_errors.Is(err, net.ErrClosed) { - break + if len(msg.Questions) != 1 { + msg.Header.Response = true + msg.Header.RCode = dnsmessage.RCodeFormatError + c.decref(msg, addr) + return + } + + opt := false + edns0 := uint16(0) + for i := range msg.Additionals { + if msg.Additionals[i].Header.Type == dnsmessage.TypeOPT { + if opt { + msg.Header.RCode = dnsmessage.RCodeFormatError + c.decref(msg, addr) + return } - continue - } - - query, err := MessageFromWireFormat(buf[:n]) - if err != nil { - errors.LogDebug(context.Background(), addr, " xdns from wireformat err ", err) - continue - } - - resp, payload := responseFor(&query, c.domains) - - var clientID [8]byte - n = copy(clientID[:], payload) - payload = payload[n:] - if n == len(clientID) { - r := bytes.NewReader(payload) - for { - p, err := nextPacketServer(r) - if err != nil { - break - } - - buf := make([]byte, len(p)) - copy(buf, p) - select { - case c.readQueue <- &packet{ - p: buf, - addr: clientIDToAddr(clientID), - }: - default: - errors.LogDebug(context.Background(), addr, " ", clientID, " mask read err queue full") - } - } - } else { - if resp != nil && resp.Rcode() == RcodeNoError { - resp.Flags |= RcodeNameError - } - } - - if resp != nil { - select { - case c.ch <- &record{resp, addr, clientIDToAddr(clientID)}: - default: - errors.LogDebug(context.Background(), addr, " ", clientID, " mask read err record queue full") + opt = true + edns0 = uint16(msg.Additionals[i].Header.Class) + if ver := (msg.Additionals[i].Header.TTL >> 16) & 0xFF; ver != 0 { + msg.Header.RCode = dnsmessage.RCodeSuccess + msg.Additionals[i].Header.TTL = 1 << 24 + c.decref(msg, addr) + return } } } + if opt { + if edns0 < 512 { + edns0 = 512 + } + if edns0 > 4096 { + edns0 = 4096 + } + } + errors.LogDebug(context.Background(), addr, " edns0 ", edns0, " buf ", len(buf), " ", msg.Questions[0].Type) - errors.LogDebug(context.Background(), "xdns closed") + var domain *Domain + for i := range c.domains { + if c.domains[i].IsDomain(msg.Questions[0].Name) { + domain = c.domains[i] + break + } + } + if domain == nil { + msg.Header.Response = true + msg.Header.RCode = dnsmessage.RCodeNameError + c.decref(msg, addr) + return + } + if !domain.HasType(uint16(msg.Questions[0].Type)) { + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + return + } - close(c.ch) - close(c.readQueue) + var decoded [255]byte + n := domain.Decode(&decoded, msg.Questions[0].Name) + if n < 9 { + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + return + } + if TypeMap_[decoded[0]&3] != uint16(msg.Questions[0].Type) || (decoded[8]&0x3F != 3 && decoded[8]&0x3F != 8) || (decoded[8]&0x3F == 3 && n < 9+3+1) || (decoded[8]&0x3F == 8 && n != 9+8) { + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + return + } + clientID := ClientIDFromRaw([8]byte(decoded[:8])) - c.mutex.Lock() - defer c.mutex.Unlock() + r := NewResp(msg, domain, edns0) + if r == nil { + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + return + } + select { + case c.recCh <- &Rec{resp: r, clientID: clientID, addr: addr}: + default: + msg.Header.Response = true + msg.Header.Authoritative = true + msg.Header.RCode = dnsmessage.RCodeSuccess + c.decref(msg, addr) + } - c.closed = true - for key, q := range c.writeQueueMap { - close(q.queue) - close(q.stash) - delete(c.writeQueueMap, key) + if decoded[8]&0x3F == 8 { + return + } + p := pool4K.Get().([]byte) + p = p[:0] + if decoded[8]&0xC0 == 0xC0 { + out := pool4K.Get().([]byte) + n := c.fragManager.Feed(out, FragKey{clientID: clientID, fragID: decoded[12]}, decoded[13], decoded[14], decoded[15:n]) + pool4K.Put(p[:cap(p)]) + if n > 0 { + p = out[:n] + } else { + pool4K.Put(out[:cap(out)]) + return + } + } else { + p = append(p, decoded[12:n]...) + } + select { + case <-c.closeCh: + pool4K.Put(p[:cap(p)]) + return + case c.readCh <- packet{p: p, addr: clientID.Addr()}: + return } } -func (c *xdnsConnServer) sendLoop() { - var nextRec *record +func (c *xdnsServer) run() { + c.wg.Add(1) + go c.recv() + + c.wg.Add(1) + go c.send() + + c.wg.Add(1) + go c.dr() + + c.wg.Wait() + close(c.readCh) + close(c.recCh) + close(c.drCh) + c.fragManager.Close() + c.sendManager.Close() +} + +func (c *xdnsServer) recv() { + defer c.wg.Done() + + var buf [512]byte + for { + n, addr, err := c.PacketConn.ReadFrom(buf[:]) + if err != nil { + if c.closed() { + return + } + errors.LogErrorInner(context.Background(), err, "recv err") + return + } + c.read(buf[:n], addr) + } +} + +func (c *xdnsServer) send() { + defer c.wg.Done() + + timer := time.NewTimer(maxResponseDelay) + timer.Stop() + var buf [4096]byte + var data [4096]byte + var nextRec *Rec for { - var err error rec := nextRec nextRec = nil if rec == nil { - var ok bool - rec, ok = <-c.ch - if !ok { - break + select { + case rec = <-c.recCh: + case <-c.closeCh: + return } } - if rec.Resp.Rcode() == RcodeNoError && len(rec.Resp.Question) == 1 { - var payload bytes.Buffer - limit := maxEncodedPayloadForType(rec.Resp.Question[0].Type) - timer := time.NewTimer(maxResponseDelay) - - for { - c.mutex.Lock() - q := c.ensureQueue(rec.ClientAddr) - if q == nil { - c.mutex.Unlock() - return - } - q.rrType = rec.Resp.Question[0].Type - c.mutex.Unlock() - - var p []byte - + ch, stash := c.sendManager.Pop(rec.clientID) + left := rec.resp.cap + timer.Reset(maxResponseDelay) + var ps [][]byte + for { + var p []byte + select { + case p = <-stash: + default: select { - case p = <-q.stash: + case p = <-stash: + case p = <-ch: default: select { - case p = <-q.stash: - case p = <-q.queue: - default: - select { - case p = <-q.stash: - case p = <-q.queue: - case <-timer.C: - case nextRec = <-c.ch: - } + case p = <-stash: + case p = <-ch: + case <-timer.C: + case nextRec = <-c.recCh: } } - - timer.Reset(0) - - if len(p) == 0 { + } + if len(p) == 0 { + break + } + timer.Reset(0) + left -= 2 + len(p) + if left < 0 { + if len(ps) == 0 { + errors.LogError(context.Background(), "err size ", len(p)) break } - - limit -= 2 + len(p) - if limit < 0 { - if payload.Len() == 0 { - errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns payload too large for rrtype ", rec.Resp.Question[0].Type, " ", len(p)) - continue - } - c.stash(q, p) - break - } - - // if len(p) > 65535 { - // panic(len(p)) - // } - - _ = binary.Write(&payload, binary.BigEndian, uint16(len(p))) - payload.Write(p) + c.sendManager.Stash(rec.clientID, p) + break } + ps = append(ps, p) + } + timer.Stop() - timer.Stop() - rec.Resp.Answer, err = answersForPayload(rec.Resp.Question[0], responseTTL, payload.Bytes()) - if err != nil { - errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns encode err ", err) - continue + d := data[:0] + for i := range ps { + l := len(ps[i]) + if i == len(ps)-1 { + l |= 0xC000 } + d = append(d, []byte{byte(l >> 8), byte(l)}...) + d = append(d, ps[i]...) } + _, _ = c.PacketConn.WriteTo(rec.resp.Encode(buf[:0], d), rec.addr) + } +} - buf, err := rec.Resp.WireFormat() - if err != nil { - errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns wireformat err ", err) - continue - } +func (c *xdnsServer) dr() { + defer c.wg.Done() - if len(buf) > maxUDPPayload { - errors.LogDebug(context.Background(), rec.Addr, " ", rec.ClientAddr, " xdns truncate ", len(buf)) - buf = buf[:maxUDPPayload] - buf[2] |= 0x02 - } - - if c.closed { + var buf [512]byte + for { + select { + case <-c.closeCh: return - } - - _, err = c.PacketConn.WriteTo(buf, rec.Addr) - if go_errors.Is(err, net.ErrClosed) { - c.closed = true - break + case r := <-c.drCh: + _, _ = c.PacketConn.WriteTo(common.Must2(r.msg.AppendPack(buf[:0])), r.addr) } } } -func (c *xdnsConnServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) { - packet, ok := <-c.readQueue - if !ok { - return 0, nil, net.ErrClosed +func (c *xdnsServer) ReadFrom(p []byte) (n int, addr net.Addr, err error) { + packet, ok := <-c.readCh + if ok { + n = copy(p, packet.p) + pool4K.Put(packet.p[:cap(packet.p)]) + return n, packet.addr, nil } - if len(p) < len(packet.p) { - errors.LogDebug(context.Background(), packet.addr, " mask read err short buffer ", len(p), " ", len(packet.p)) - return 0, packet.addr, nil - } - copy(p, packet.p) - return len(packet.p), packet.addr, nil + return 0, nil, io.ErrClosedPipe } -func (c *xdnsConnServer) WriteTo(p []byte, addr net.Addr) (n int, err error) { - c.mutex.Lock() - defer c.mutex.Unlock() - - q := c.ensureQueue(addr) - if q == nil { +func (c *xdnsServer) WriteTo(p []byte, addr net.Addr) (n int, err error) { + if c.closed() { return 0, io.ErrClosedPipe } - limit := maxEncodedPayloadForType(q.rrType) - if q.rrType == 0 { - limit = maxEncodedPayloadTXT - } - if len(p)+2 > limit { - errors.LogDebug(context.Background(), addr, " mask write err short write ", len(p), "+2 > ", limit) - return 0, nil - } - - buf := make([]byte, len(p)) - copy(buf, p) - - select { - case q.queue <- buf: - return len(p), nil - default: - // errors.LogDebug(context.Background(), addr, " mask write err queue full") - return 0, nil + if len(p) == 0 || len(p) > 4096 { + errors.LogError(context.Background(), "err size ", len(p)) + return 0, errors.New("err size") } + c.sendManager.Push(ClientIDFromAddr(addr.(*net.UDPAddr)), p) + return len(p), nil } -func (c *xdnsConnServer) Close() error { - c.closed = true - return c.PacketConn.Close() +func (c *xdnsServer) Close() error { + c.mu.Lock() + defer c.mu.Unlock() + if c.closed() { + return nil + } + close(c.closeCh) + _ = c.PacketConn.Close() + return nil } -func nextPacketServer(r *bytes.Reader) ([]byte, error) { - eof := func(err error) error { - if err == io.EOF { - err = io.ErrUnexpectedEOF - } - return err - } +func (c *xdnsServer) SetDeadline(t time.Time) error { return errors.New("not support") } - for { - prefix, err := r.ReadByte() - if err != nil { - return nil, err - } - if prefix >= 224 { - paddingLen := prefix - 224 - _, err := io.CopyN(io.Discard, r, int64(paddingLen)) - if err != nil { - return nil, eof(err) - } - } else { - p := make([]byte, int(prefix)) - _, err = io.ReadFull(r, p) - return p, eof(err) - } - } -} +func (c *xdnsServer) SetReadDeadline(t time.Time) error { return errors.New("not support") } -func responseFor(query *Message, domains []domainSpec) (*Message, []byte) { - resp := &Message{ - ID: query.ID, - Flags: 0x8000, - Question: query.Question, - } - - if query.Flags&0x8000 != 0 { - return nil, nil - } - - payloadSize := 0 - for _, rr := range query.Additional { - if rr.Type != RRTypeOPT { - continue - } - if len(resp.Additional) != 0 { - resp.Flags |= RcodeFormatError - return resp, nil - } - resp.Additional = append(resp.Additional, RR{ - Name: Name{}, - Type: RRTypeOPT, - Class: 4096, - TTL: 0, - Data: []byte{}, - }) - additional := &resp.Additional[0] - - version := (rr.TTL >> 16) & 0xff - if version != 0 { - resp.Flags |= ExtendedRcodeBadVers & 0xf - additional.TTL = (ExtendedRcodeBadVers >> 4) << 24 - return resp, nil - } - - payloadSize = int(rr.Class) - } - if payloadSize < 512 { - payloadSize = 512 - } - - if len(query.Question) != 1 { - resp.Flags |= RcodeFormatError - return resp, nil - } - question := query.Question[0] - - var ( - prefix Name - ok bool - match domainSpec - ) - for _, domain := range domains { - prefix, ok = question.Name.TrimSuffix(domain.name) - if ok { - match = domain - break - } - } - if !ok { - resp.Flags |= RcodeNameError - return resp, nil - } - resp.Flags |= 0x0400 - - if query.Opcode() != 0 { - resp.Flags |= RcodeNotImplemented - return resp, nil - } - - switch question.Type { - case RRTypeTXT, RRTypeA, RRTypeAAAA: - default: - resp.Flags |= RcodeNameError - return resp, nil - } - if match.rrType != 0 && question.Type != match.rrType { - resp.Flags |= RcodeNameError - return resp, nil - } - - encoded := bytes.ToUpper(bytes.Join(prefix, nil)) - payload := make([]byte, base32Encoding.DecodedLen(len(encoded))) - n, err := base32Encoding.Decode(payload, encoded) - if err != nil { - resp.Flags |= RcodeNameError - return resp, nil - } - payload = payload[:n] - - if payloadSize < maxUDPPayload { - resp.Flags |= RcodeFormatError - return resp, nil - } - - return resp, payload -} +func (c *xdnsServer) SetWriteDeadline(t time.Time) error { return errors.New("not support") } diff --git a/transport/internet/finalmask/xdns/spec.go b/transport/internet/finalmask/xdns/spec.go deleted file mode 100644 index 28461569f..000000000 --- a/transport/internet/finalmask/xdns/spec.go +++ /dev/null @@ -1,80 +0,0 @@ -package xdns - -import ( - "strings" - - "github.com/xtls/xray-core/common/errors" -) - -type domainSpec struct { - name Name - rrType uint16 -} - -func rrTypeFromMethod(method string) (uint16, error) { - switch strings.ToLower(method) { - case "", "txt": - return RRTypeTXT, nil - case "a": - return RRTypeA, nil - case "aaaa": - return RRTypeAAAA, nil - default: - return 0, errors.New("unsupported method") - } -} - -func parseDomainSpec(s string, defaultMethod string) (domainSpec, error) { - domainPart := s - method := "" - hasMethod := false - - if i := strings.LastIndex(s, ":"); i >= 0 { - domainPart = s[:i] - method = s[i+1:] - hasMethod = true - } else if defaultMethod != "" { - method = defaultMethod - hasMethod = true - } - - if domainPart == "" { - return domainSpec{}, errors.New("empty domain") - } - - name, err := ParseName(domainPart) - if err != nil { - return domainSpec{}, err - } - - rrType := uint16(0) - if hasMethod { - var err error - rrType, err = rrTypeFromMethod(method) - if err != nil { - return domainSpec{}, err - } - } - - return domainSpec{ - name: name, - rrType: rrType, - }, nil -} - -func parseResolver(s string) (Name, string, uint16, error) { - head, server, ok := strings.Cut(s, "+udp://") - if !ok { - return nil, "", 0, errors.New("invalid resolver scheme") - } - if server == "" { - return nil, "", 0, errors.New("empty resolver server") - } - - spec, err := parseDomainSpec(head, "txt") - if err != nil { - return nil, "", 0, err - } - - return spec.name, server, spec.rrType, nil -} diff --git a/transport/internet/finalmask/xdns/xdns_test.go b/transport/internet/finalmask/xdns/xdns_test.go new file mode 100644 index 000000000..d2a4d1c1b --- /dev/null +++ b/transport/internet/finalmask/xdns/xdns_test.go @@ -0,0 +1,208 @@ +package xdns + +import ( + "bytes" + "crypto/rand" + "fmt" + mrand "math/rand" + "testing" + + "github.com/xtls/xray-core/common" + "golang.org/x/net/dns/dnsmessage" +) + +func TestXxx(t *testing.T) { + m1 := dnsmessage.Message{ + Questions: []dnsmessage.Question{ + { + Name: dnsmessage.MustNewName("a.example.com."), + }, + }, + Answers: []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeA, + Class: dnsmessage.ClassINET, + TTL: 60, + Length: 16, + }, + Body: &dnsmessage.AResource{A: [4]byte{127, 0, 0, 1}}, + }, + }, + Additionals: []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("."), + Type: dnsmessage.TypeOPT, + Class: 255, + TTL: 0, + Length: 16, + }, + Body: &dnsmessage.OPTResource{}, + }, + }, + } + p1, e1 := m1.Pack() + if e1 != nil { + t.Fatal(e1) + } + if !bytes.Equal(p1, []byte{ + 0, 0, 0, 0, 0, 1, 0, 1, 0, 0, 0, 1, + 1, 97, 7, 101, 120, 97, 109, 112, 108, 101, 3, 99, 111, 109, 0, + 0, 0, + 0, 0, + 192, 12, + 0, 1, + 0, 1, + 0, 0, 0, 60, + 0, 4, + 127, 0, 0, 1, + 0, + 0, 41, + 0, 255, + 0, 0, 0, 0, + 0, 0, + }) { + t.Fatal("!bytes.Equal") + } + + domain, _ := NewDomain("a.example.com", 200, 1, []uint16{1}, 0) + fmt.Println(domain.cap, domain.lenMax) + lenMax := domain.lenMax + data := make([]byte, domain.cap) + msg := dnsmessage.Message{} + msg.Unpack(p1) + for range 3 { + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + n := mrand.Intn(255) + for range n { + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeA, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.AResource{A: [4]byte{127, 0, 0, 1}}, + }) + } + if len(common.Must2(msg.Pack())) != 12+15+2+2+n*(2+2+2+4+2+4) { + t.Fatal("fatal a") + } + } + for range 3 { + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + n := mrand.Intn(255) + for range n { + common.Must2(rand.Read(data)) + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeCNAME, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.CNAMEResource{ + CNAME: domain.Encode(data), + }, + }) + } + if len(common.Must2(msg.Pack())) > 12+15+2+2+n*(2+2+2+4+2+lenMax) { + t.Fatal("fatal cname") + } + } + for range 3 { + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + n := (mrand.Intn(2048) + 1024) % 2048 + a := n / 255 + b := n % 255 + c := 0 + var d [255]byte + var s []string + for range a { + s = append(s, string(d[:])) + } + if b > 0 { + c = 1 + s = append(s, string(d[:b])) + } + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeTXT, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.TXTResource{TXT: s}, + }) + if len(common.Must2(msg.Pack())) != 12+15+2+2+(2+2+2+4+2+n+n/255+c) { + t.Fatal("fatal txt") + } + } + for range 3 { + msg.Answers = nil + msg.Authorities = nil + msg.Additionals = nil + n := mrand.Intn(255) + for range n { + msg.Answers = append(msg.Answers, dnsmessage.Resource{ + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("a.example.com."), + Type: dnsmessage.TypeAAAA, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.AAAAResource{AAAA: [16]byte{}}, + }) + } + if len(common.Must2(msg.Pack())) != 12+15+2+2+n*(2+2+2+4+2+16) { + t.Fatal("fatal aaaa") + } + } +} + +func TestTXT(t *testing.T) { + txt := [][]byte{{}, {}} + for i := range 255 { + txt[0] = append(txt[0], byte(i)) + } + txt[1] = []byte{255} + str := []string{} + for i := range txt { + str = append(str, string(txt[i])) + } + m1 := dnsmessage.Message{ + Answers: []dnsmessage.Resource{ + { + Header: dnsmessage.ResourceHeader{ + Name: dnsmessage.MustNewName("."), + Type: dnsmessage.TypeTXT, + Class: dnsmessage.ClassINET, + TTL: 60, + }, + Body: &dnsmessage.TXTResource{ + TXT: str, + }, + }, + }, + } + p1 := common.Must2(m1.Pack()) + + m2 := dnsmessage.Message{} + common.Must(m2.Unpack(p1)) + if len(m2.Answers[0].Body.(*dnsmessage.TXTResource).TXT) != len(txt) { + t.Fatal("fatal txt") + } + for i := range txt { + if !bytes.Equal(txt[i], []byte(m2.Answers[0].Body.(*dnsmessage.TXTResource).TXT[i])) { + t.Fatal("fatal txt") + } + } +} diff --git a/transport/internet/finalmask/xicmp/client.go b/transport/internet/finalmask/xicmp/client.go index cabdb1942..702eefb9b 100644 --- a/transport/internet/finalmask/xicmp/client.go +++ b/transport/internet/finalmask/xicmp/client.go @@ -310,13 +310,6 @@ func (c *xicmpConnClient) Close() error { _ = c.icmp4.Close() _ = c.icmp6.Close() c.wg.Wait() - select { - case p := <-c.readCh: - if p.p != nil { - pool.Put(p.p) - } - default: - } close(c.readCh) return nil } diff --git a/transport/internet/finalmask/xicmp/server.go b/transport/internet/finalmask/xicmp/server.go index 117638923..1d99d03b7 100644 --- a/transport/internet/finalmask/xicmp/server.go +++ b/transport/internet/finalmask/xicmp/server.go @@ -329,13 +329,6 @@ func (c *xicmpConnServer) Close() error { _ = c.icmp4.Close() _ = c.icmp6.Close() c.wg.Wait() - select { - case p := <-c.readCh: - if p.p != nil { - pool.Put(p.p) - } - default: - } close(c.readCh) return nil } diff --git a/transport/internet/finalmask/xicmp/server_oob.go b/transport/internet/finalmask/xicmp/server_oob.go index 8d1c26db8..8eebf43bb 100644 --- a/transport/internet/finalmask/xicmp/server_oob.go +++ b/transport/internet/finalmask/xicmp/server_oob.go @@ -340,13 +340,6 @@ func (c *xicmpConnServer) Close() error { _ = c.icmp4.Close() _ = c.icmp6.Close() c.wg.Wait() - select { - case p := <-c.readCh: - if p.p != nil { - pool.Put(p.p) - } - default: - } close(c.readCh) return nil } diff --git a/transport/internet/httpupgrade/dialer.go b/transport/internet/httpupgrade/dialer.go index d05ae8f47..cf244b38d 100644 --- a/transport/internet/httpupgrade/dialer.go +++ b/transport/internet/httpupgrade/dialer.go @@ -3,6 +3,8 @@ package httpupgrade import ( "bufio" "context" + "crypto/rand" + "encoding/base64" "net/http" "net/url" "strings" @@ -97,6 +99,16 @@ func dialhttpUpgrade(ctx context.Context, dest net.Destination, streamSettings * req.Header.Set("Connection", "Upgrade") req.Header.Set("Upgrade", "websocket") + // make a valid Sec-WebSocket-Key if not present + if len(req.Header.Values("Sec-WebSocket-Key")) == 0 { + var buf [16]byte + rand.Read(buf[:]) + req.Header.Set("Sec-WebSocket-Key", base64.StdEncoding.EncodeToString(buf[:])) + } + if len(req.Header.Values("Sec-WebSocket-Version")) == 0 { + req.Header.Set("Sec-WebSocket-Version", "13") + } + err = req.Write(conn) if err != nil { return nil, err diff --git a/transport/internet/httpupgrade/hub.go b/transport/internet/httpupgrade/hub.go index 9a6429447..2e968840a 100644 --- a/transport/internet/httpupgrade/hub.go +++ b/transport/internet/httpupgrade/hub.go @@ -3,7 +3,9 @@ package httpupgrade import ( "bufio" "context" + "crypto/sha1" "crypto/tls" + "encoding/base64" "io" "net/http" "strings" @@ -81,6 +83,11 @@ func (s *server) upgrade(conn net.Conn) (stat.Connection, error) { } resp.Header.Set("Connection", "Upgrade") resp.Header.Set("Upgrade", "websocket") + // respond a valid Sec-WebSocket-Accept header if received a Sec-WebSocket-Key + if wsKey := req.Header.Get("Sec-WebSocket-Key"); wsKey != "" { + acceptKey := sha1.Sum([]byte(wsKey + "258EAFA5-E914-47DA-95CA-C5AB0DC85B11")) // magic number in RFC 6455 + resp.Header.Set("Sec-WebSocket-Accept", base64.StdEncoding.EncodeToString(acceptKey[:])) + } err = resp.Write(conn) if err != nil { return nil, err diff --git a/transport/internet/hysteria/dialer.go b/transport/internet/hysteria/dialer.go index f37cdd11c..4b774c353 100644 --- a/transport/internet/hysteria/dialer.go +++ b/transport/internet/hysteria/dialer.go @@ -119,7 +119,7 @@ func (c *client) dial(ctx context.Context) error { if err != nil { return errors.New("failed to dial to dest").Base(err) } - pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn + pktConn = conn.(*net.PacketConnWrapper).PacketConn udpAddr = conn.RemoteAddr() } else { conn, err := internet.DialSystem(ctx, c.dest, c.socketConfig) @@ -127,7 +127,7 @@ func (c *client) dial(ctx context.Context) error { return errors.New("failed to dial to dest").Base(err) } switch c := conn.(type) { - case *internet.PacketConnWrapper: + case *net.PacketConnWrapper: pktConn = c.PacketConn udpAddr = c.RemoteAddr() case *cnc.Connection: diff --git a/transport/internet/kcp/dialer.go b/transport/internet/kcp/dialer.go index 8ec4d1973..443befc73 100644 --- a/transport/internet/kcp/dialer.go +++ b/transport/internet/kcp/dialer.go @@ -57,7 +57,7 @@ func DialKCP(ctx context.Context, dest net.Destination, streamSettings *internet conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) } if err != nil { - return nil, errors.New("failed to dial to dest: ", err).AtWarning().Base(err) + return nil, errors.New("failed to dial to dest: ", err).Base(err) } kcpSettings := streamSettings.ProtocolSettings.(*Config) diff --git a/transport/internet/masque/conn.go b/transport/internet/masque/conn.go index 6aeaf162b..9ccea3d83 100644 --- a/transport/internet/masque/conn.go +++ b/transport/internet/masque/conn.go @@ -23,9 +23,23 @@ func (e *PacketTooBigError) Error() string { return "packet too big for the tunnel" } +type httpConn interface { + LocalAddr() net.Addr + RemoteAddr() net.Addr + Close() error +} + +type quicConn struct { + *quic.Conn +} + +func (c quicConn) Close() error { + return c.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "") +} + type Conn struct { ipConn *connectip.Conn - quicConn *quic.Conn + httpConn httpConn local []netip.Addr closeOnce sync.Once } @@ -58,17 +72,17 @@ func (c *Conn) Write(b []byte) (int, error) { func (c *Conn) Close() error { c.closeOnce.Do(func() { c.ipConn.Close() - c.quicConn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "") + c.httpConn.Close() }) return nil } func (c *Conn) LocalAddr() net.Addr { - return c.quicConn.LocalAddr() + return c.httpConn.LocalAddr() } func (c *Conn) RemoteAddr() net.Addr { - return c.quicConn.RemoteAddr() + return c.httpConn.RemoteAddr() } func (c *Conn) SetDeadline(time.Time) error { diff --git a/transport/internet/masque/connectip/conn.go b/transport/internet/masque/connectip/conn.go index 3e5fc6362..510c8e567 100644 --- a/transport/internet/masque/connectip/conn.go +++ b/transport/internet/masque/connectip/conn.go @@ -38,22 +38,30 @@ const ( ipProtoICMPv6 = 58 ) -type http3Stream interface { +type requestStream interface { io.ReadWriteCloser - StreamID() quic.StreamID - ReceiveDatagram(context.Context) ([]byte, error) - SendDatagram([]byte) error CancelRead(quic.StreamErrorCode) CancelWrite(quic.StreamErrorCode) SetWriteDeadline(time.Time) error } +type http3Stream interface { + requestStream + StreamID() quic.StreamID + ReceiveDatagram(context.Context) ([]byte, error) + SendDatagram([]byte) error +} + var ( _ http3Stream = &http3.Stream{} _ http3Stream = &http3.RequestStream{} ) -const maxQueuedCapsules = 128 +const ( + maxQueuedCapsules = 128 + maxQueuedDatagrams = 128 + maxCapsulePacketSize = 1<<16 - 1 +) var errCapsuleLimit = goerrors.New("connect-ip: capsule limit exceeded") @@ -63,7 +71,10 @@ type streamWrite struct { } type Conn struct { - str http3Stream + str requestStream + h3 http3Stream + datagrams chan []byte + writeMu sync.Mutex writeNotify chan struct{} writeDone chan error @@ -87,7 +98,7 @@ type Conn struct { datagramCapsuleOnce sync.Once } -func newProxiedConn(str http3Stream) *Conn { +func newProxiedConn(str requestStream) *Conn { c := &Conn{ str: str, writeNotify: make(chan struct{}, 1), @@ -97,6 +108,9 @@ func newProxiedConn(str http3Stream) *Conn { availableRouteUpdates: make(chan []IPRoute, 1), closeChan: make(chan struct{}), } + if c.h3, _ = str.(http3Stream); c.h3 == nil { + c.datagrams = make(chan []byte, maxQueuedDatagrams) + } go func() { err := c.readFromStream() c.mu.Lock() @@ -382,6 +396,12 @@ func (c *Conn) readFromStream() error { } queueLatest(c.availableRouteUpdates, capsule.IPAddressRanges) case capsuleTypeDatagram: + if c.h3 == nil { + if err := c.queueDatagram(cr); err != nil { + return err + } + continue + } c.datagramCapsuleOnce.Do(func() { errors.LogWarning(context.Background(), "connect-ip: dropping IP packets sent in DATAGRAM capsules, only QUIC DATAGRAM frames are supported") }) @@ -412,7 +432,7 @@ func (c *Conn) writeToStream() error { if w.Fin { return c.str.Close() } - if _, err := c.str.Write(w.Data); err != nil { + if err := c.write(w.Data); err != nil { return err } } @@ -420,6 +440,41 @@ func (c *Conn) writeToStream() error { return c.closeErr } +func (c *Conn) write(b []byte) error { + c.writeMu.Lock() + defer c.writeMu.Unlock() + _, err := c.str.Write(b) + return err +} + +func (c *Conn) queueDatagram(cr http3.CapsuleReader) error { + if cr.Remaining() > int64(len(contextIDZero)+maxCapsulePacketSize) { + errors.LogDebug(context.Background(), "connect-ip: dropping a ", cr.Remaining(), "-byte DATAGRAM capsule") + return cr.Discard() + } + data := make([]byte, cr.Remaining()) + if _, err := io.ReadFull(cr, data); err != nil { + return err + } + select { + case c.datagrams <- data: + case <-c.closeChan: + } + return nil +} + +func (c *Conn) receiveDatagram() ([]byte, error) { + if c.h3 != nil { + return c.h3.ReceiveDatagram(context.Background()) + } + select { + case data := <-c.datagrams: + return data, nil + case <-c.closeChan: + return nil, c.closeErr + } +} + func (c *Conn) ReadPacket(b []byte) (int, error) { for { select { @@ -427,7 +482,7 @@ func (c *Conn) ReadPacket(b []byte) (int, error) { return 0, c.closeErr default: } - data, err := c.str.ReceiveDatagram(context.Background()) + data, err := c.receiveDatagram() if err != nil { select { case <-c.closeChan: @@ -525,7 +580,18 @@ func (c *Conn) WritePacket(b []byte) (icmp []byte, err error) { errors.LogDebugInner(context.Background(), err, "dropping proxied packet (", len(b), " bytes) that can't be proxied") return nil, nil } - if err := c.str.SendDatagram(data); err != nil { + if c.h3 == nil { + if err := c.write(data); err != nil { + select { + case <-c.closeChan: + return nil, c.closeErr + default: + return nil, err + } + } + return nil, nil + } + if err := c.h3.SendDatagram(data); err != nil { if tooLarge, ok := goerrors.AsType[*quic.DatagramTooLargeError](err); ok { icmpPacket, err := composeICMPTooLargePacket(b, int(tooLarge.MaxDatagramPayloadSize)-c.datagramOverhead()) if err != nil { @@ -578,14 +644,22 @@ func (c *Conn) composeDatagram(b []byte) ([]byte, error) { } b[7]-- } - data := make([]byte, 0, len(contextIDZero)+len(b)) + size := len(contextIDZero) + len(b) + var data []byte + if c.h3 == nil { + data = make([]byte, 0, quicvarint.Len(uint64(capsuleTypeDatagram))+quicvarint.Len(uint64(size))+size) + data = quicvarint.Append(data, uint64(capsuleTypeDatagram)) + data = quicvarint.Append(data, uint64(size)) + } else { + data = make([]byte, 0, size) + } data = append(data, contextIDZero...) data = append(data, b...) return data, nil } func (c *Conn) datagramOverhead() int { - return quicvarint.Len(uint64(c.str.StreamID()/4)) + len(contextIDZero) + return quicvarint.Len(uint64(c.h3.StreamID()/4)) + len(contextIDZero) } func (c *Conn) MaxPacketSize() int { @@ -594,7 +668,10 @@ func (c *Conn) MaxPacketSize() int { return 0 default: } - err := c.str.SendDatagram(make([]byte, 1<<16)) + if c.h3 == nil { + return maxCapsulePacketSize + } + err := c.h3.SendDatagram(make([]byte, 1<<16)) tooLarge, ok := goerrors.AsType[*quic.DatagramTooLargeError](err) if !ok { return 0 diff --git a/transport/internet/masque/connectip/http2.go b/transport/internet/masque/connectip/http2.go new file mode 100644 index 000000000..be4b52d72 --- /dev/null +++ b/transport/internet/masque/connectip/http2.go @@ -0,0 +1,198 @@ +package connectip + +import ( + "bufio" + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "os" + "sync" + "time" + + "github.com/apernet/quic-go" +) + +const maxStreamBuffer = 32 << 10 + +type HTTP2ClientConn struct { + roundTripper http.RoundTripper +} + +func NewHTTP2ClientConn(rt http.RoundTripper) *HTTP2ClientConn { + return &HTTP2ClientConn{roundTripper: rt} +} + +func (c *HTTP2ClientConn) Dial(req *Request) (*Conn, *http.Response, error) { + httpReq := req.httpRequest() + if httpReq.URL == nil { + return nil, nil, errors.New("connect-ip: request URL is nil") + } + if httpReq.Host == "" && httpReq.URL.Host == "" { + return nil, nil, errors.New("connect-ip: request needs a host") + } + + ctx := httpReq.Context() + streamCtx, cancel := context.WithCancel(context.WithoutCancel(ctx)) + stop := context.AfterFunc(ctx, cancel) + body := NewStreamBuffer() + r := httpReq.Clone(streamCtx) + r.Header[":protocol"] = []string{requestProtocol} + r.Body = body + rsp, err := c.roundTripper.RoundTrip(r) + if !stop() { + if err == nil { + rsp.Body.Close() + } + err = context.Cause(ctx) + } + if err != nil { + cancel() + return nil, nil, fmt.Errorf("connect-ip: failed to send request: %w", err) + } + if rsp.StatusCode < 200 || rsp.StatusCode > 299 { + cancel() + rsp.Body.Close() + return nil, rsp, fmt.Errorf("connect-ip: server responded with %d", rsp.StatusCode) + } + return newProxiedConn(&http2Stream{ + reader: bufio.NewReader(rsp.Body), + body: body, + rsp: rsp.Body, + cancel: cancel, + }), rsp, nil +} + +type http2Stream struct { + reader *bufio.Reader + body *StreamBuffer + rsp io.Closer + cancel context.CancelFunc +} + +func (s *http2Stream) Read(b []byte) (int, error) { return s.reader.Read(b) } +func (s *http2Stream) ReadByte() (byte, error) { return s.reader.ReadByte() } +func (s *http2Stream) Write(b []byte) (int, error) { return s.body.Write(b) } +func (s *http2Stream) Close() error { return s.body.Close() } +func (s *http2Stream) CancelRead(quic.StreamErrorCode) { s.abort() } +func (s *http2Stream) CancelWrite(quic.StreamErrorCode) { s.abort() } +func (s *http2Stream) SetWriteDeadline(t time.Time) error { return s.body.SetWriteDeadline(t) } + +func (s *http2Stream) abort() { + s.cancel() + s.body.CloseWithError(net.ErrClosed) + s.rsp.Close() +} + +type StreamBuffer struct { + mu sync.Mutex + cond sync.Cond + buf []byte + closed bool + err error + deadline time.Time +} + +func NewStreamBuffer() *StreamBuffer { + b := &StreamBuffer{} + b.cond.L = &b.mu + return b +} + +func (b *StreamBuffer) Read(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + for len(b.buf) == 0 && !b.closed && b.err == nil { + b.cond.Wait() + } + if b.err != nil { + return 0, b.err + } + if len(b.buf) == 0 { + return 0, io.EOF + } + n := copy(p, b.buf) + b.buf = b.buf[:copy(b.buf, b.buf[n:])] + b.cond.Broadcast() + return n, nil +} + +func (b *StreamBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + for { + switch { + case b.err != nil: + return 0, b.err + case b.closed: + return 0, io.ErrClosedPipe + case !b.deadline.IsZero() && !time.Now().Before(b.deadline): + return 0, os.ErrDeadlineExceeded + case len(b.buf) < maxStreamBuffer: + b.buf = append(b.buf, p...) + b.cond.Broadcast() + return len(p), nil + } + b.cond.Wait() + } +} + +func (b *StreamBuffer) Close() error { + b.mu.Lock() + b.closed = true + b.cond.Broadcast() + b.mu.Unlock() + return nil +} + +func (b *StreamBuffer) CloseWithError(err error) { + b.mu.Lock() + if b.err == nil { + b.err = err + b.buf = nil + } + b.cond.Broadcast() + b.mu.Unlock() +} + +func (b *StreamBuffer) SetWriteDeadline(t time.Time) error { + b.mu.Lock() + b.deadline = t + b.cond.Broadcast() + b.mu.Unlock() + if d := time.Until(t); d > 0 { + time.AfterFunc(d, func() { + b.mu.Lock() + b.cond.Broadcast() + b.mu.Unlock() + }) + } + return nil +} + +type http2ResponseStream struct { + reader *bufio.Reader + body io.Closer + w io.Writer + controller *http.ResponseController +} + +func (s *http2ResponseStream) Read(b []byte) (int, error) { return s.reader.Read(b) } +func (s *http2ResponseStream) ReadByte() (byte, error) { return s.reader.ReadByte() } + +func (s *http2ResponseStream) Write(b []byte) (int, error) { + n, err := s.w.Write(b) + if err == nil { + err = s.controller.Flush() + } + return n, err +} + +func (s *http2ResponseStream) Close() error { return nil } +func (s *http2ResponseStream) CancelRead(quic.StreamErrorCode) { s.body.Close() } +func (s *http2ResponseStream) CancelWrite(quic.StreamErrorCode) { s.body.Close() } +func (s *http2ResponseStream) SetWriteDeadline(t time.Time) error { + return s.controller.SetWriteDeadline(t) +} diff --git a/transport/internet/masque/connectip/http2_test.go b/transport/internet/masque/connectip/http2_test.go new file mode 100644 index 000000000..4ddc9c52f --- /dev/null +++ b/transport/internet/masque/connectip/http2_test.go @@ -0,0 +1,502 @@ +package connectip + +import ( + "bufio" + "bytes" + "context" + "errors" + "io" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "os" + "slices" + "sync" + "testing" + "time" + + "github.com/apernet/quic-go/http3" + "github.com/apernet/quic-go/quicvarint" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/net/ipv4" + "golang.org/x/net/ipv6" +) + +type roundTripFunc func(*http.Request) (*http.Response, error) + +func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) { return f(r) } + +type pipeResponseWriter struct { + *io.PipeWriter + header http.Header + status int + headerOnce sync.Once + headerDone chan struct{} +} + +func (w *pipeResponseWriter) Header() http.Header { return w.header } + +func (w *pipeResponseWriter) WriteHeader(code int) { + w.headerOnce.Do(func() { + w.status = code + close(w.headerDone) + }) +} + +func (w *pipeResponseWriter) Write(b []byte) (int, error) { + w.WriteHeader(http.StatusOK) + return w.PipeWriter.Write(b) +} + +func (w *pipeResponseWriter) Flush() {} + +func http2RoundTripper(handler http.HandlerFunc) http.RoundTripper { + return roundTripFunc(func(r *http.Request) (*http.Response, error) { + pr, pw := io.Pipe() + w := &pipeResponseWriter{PipeWriter: pw, header: http.Header{}, headerDone: make(chan struct{})} + sr := r.Clone(r.Context()) + sr.Proto, sr.ProtoMajor, sr.ProtoMinor = "HTTP/2.0", 2, 0 + go func() { + handler(w, sr) + w.WriteHeader(http.StatusOK) + pw.Close() + }() + <-w.headerDone + return &http.Response{StatusCode: w.status, Header: w.header, Body: pr}, nil + }) +} + +func setupHTTP2Conns(t *testing.T) (client, server *Conn) { + t.Helper() + + serverConns := make(chan *Conn, 1) + rt := http2RoundTripper(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, "Bearer token", r.Header.Get("Authorization")) + req, err := ParseProxyRequest(r) + if !assert.NoError(t, err) { + w.WriteHeader(http.StatusBadRequest) + return + } + conn, err := (&Proxy{}).Proxy(w, req) + if !assert.NoError(t, err) { + return + } + serverConns <- conn + <-conn.closeChan + }) + + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + req, err := NewRequest(ctx, "https://example.org/connect-ip") + require.NoError(t, err) + req.Header().Set("Authorization", "Bearer token") + client, rsp, err := NewHTTP2ClientConn(rt).Dial(req) + require.NoError(t, err) + t.Cleanup(func() { client.Close() }) + require.Equal(t, http.StatusOK, rsp.StatusCode) + require.Equal(t, "?1", rsp.Header.Get("Capsule-Protocol")) + + select { + case <-time.After(5 * time.Second): + t.Fatal("timed out") + case server = <-serverConns: + } + t.Cleanup(func() { server.Close() }) + return client, server +} + +func newTestHTTP2Stream() (*http2Stream, *io.PipeWriter) { + pr, pw := io.Pipe() + return &http2Stream{reader: bufio.NewReader(pr), body: NewStreamBuffer(), rsp: pr, cancel: func() {}}, pw +} + +func TestHTTP2Request(t *testing.T) { + requests := make(chan *http.Request, 1) + pr, pw := io.Pipe() + defer pw.Close() + rt := roundTripFunc(func(r *http.Request) (*http.Response, error) { + requests <- r + return &http.Response{StatusCode: http.StatusOK, Body: pr}, nil + }) + req, err := NewRequest(t.Context(), "https://proxy.example:8443/.well-known/masque/ip/*/*/") + require.NoError(t, err) + req.Header().Set("Authorization", "Bearer token") + conn, _, err := NewHTTP2ClientConn(rt).Dial(req) + require.NoError(t, err) + defer conn.Close() + + r := <-requests + require.Equal(t, http.MethodConnect, r.Method) + require.Equal(t, []string{requestProtocol}, r.Header[":protocol"]) + require.Equal(t, "?1", r.Header.Get("Capsule-Protocol")) + require.Equal(t, "Bearer token", r.Header.Get("Authorization")) + require.Equal(t, "proxy.example:8443", r.Host) + require.Equal(t, "https", r.URL.Scheme) + require.Equal(t, "/.well-known/masque/ip/*/*/", r.URL.Path) + require.NotNil(t, r.Body) + require.Empty(t, req.Header().Values(":protocol")) + require.Equal(t, maxCapsulePacketSize, conn.MaxPacketSize()) +} + +func TestHTTP2DialErrors(t *testing.T) { + newReq := func(ctx context.Context) *Request { + req, err := NewRequest(ctx, "https://example.org/connect-ip") + require.NoError(t, err) + return req + } + + t.Run("status", func(t *testing.T) { + var streamCtx context.Context + rt := roundTripFunc(func(r *http.Request) (*http.Response, error) { + streamCtx = r.Context() + return &http.Response{StatusCode: http.StatusForbidden, Body: io.NopCloser(bytes.NewReader(nil))}, nil + }) + _, rsp, err := NewHTTP2ClientConn(rt).Dial(newReq(t.Context())) + require.EqualError(t, err, "connect-ip: server responded with 403") + require.Equal(t, http.StatusForbidden, rsp.StatusCode) + require.ErrorIs(t, streamCtx.Err(), context.Canceled) + }) + + t.Run("round trip", func(t *testing.T) { + errRoundTrip := errors.New("extended connect not supported by peer") + rt := roundTripFunc(func(*http.Request) (*http.Response, error) { return nil, errRoundTrip }) + _, _, err := NewHTTP2ClientConn(rt).Dial(newReq(t.Context())) + require.ErrorIs(t, err, errRoundTrip) + }) + + t.Run("context", func(t *testing.T) { + rt := roundTripFunc(func(r *http.Request) (*http.Response, error) { + <-r.Context().Done() + return nil, r.Context().Err() + }) + ctx, cancel := context.WithTimeout(t.Context(), 50*time.Millisecond) + defer cancel() + _, _, err := NewHTTP2ClientConn(rt).Dial(newReq(ctx)) + require.ErrorIs(t, err, context.DeadlineExceeded) + }) +} + +func TestHTTP2Packets(t *testing.T) { + client, server := setupHTTP2Conns(t) + clientV4 := netip.MustParseAddr("192.0.2.2") + clientV6 := netip.MustParseAddr("2001:db8::2") + require.NoError(t, server.AssignAddresses([]netip.Prefix{netip.PrefixFrom(clientV4, 32), netip.PrefixFrom(clientV6, 128)})) + require.NoError(t, server.AdvertiseRoute([]IPRoute{ + {StartIP: netip.IPv4Unspecified(), EndIP: netip.MustParseAddr("255.255.255.255")}, + {StartIP: netip.IPv6Unspecified(), EndIP: netip.MustParseAddr("ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff")}, + })) + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + _, err := client.ReceiveAddressAssignment(ctx) + require.NoError(t, err) + _, err = client.Routes(ctx) + require.NoError(t, err) + require.Equal(t, maxCapsulePacketSize, client.MaxPacketSize()) + + for _, tc := range []struct { + name string + up []byte + down []byte + ttlOff int + }{ + { + name: "IPv4", + up: ipv4Packet(64, 17, clientV4, testDst4, nil, []byte("foobar")), + down: ipv4Packet(64, 17, testDst4, clientV4, nil, []byte("barfoo")), + ttlOff: 8, + }, + { + name: "IPv6 larger than a QUIC datagram", + up: ipv6Packet(64, 17, clientV6, testDst6, bytes.Repeat([]byte("up"), 4500)), + down: ipv6Packet(64, 17, testDst6, clientV6, bytes.Repeat([]byte("down"), 2250)), + ttlOff: 7, + }, + } { + t.Run(tc.name, func(t *testing.T) { + for _, dir := range []struct { + from, to *Conn + packet []byte + }{ + {client, server, tc.up}, + {server, client, tc.down}, + } { + icmp, err := dir.from.WritePacket(slices.Clone(dir.packet)) + require.NoError(t, err) + require.Nil(t, icmp) + b := make([]byte, 1<<16) + n, err := dir.to.ReadPacket(b) + require.NoError(t, err) + require.Len(t, b[:n], len(dir.packet)) + require.Equal(t, dir.packet[tc.ttlOff]-1, b[tc.ttlOff]) + if tc.ttlOff == 8 { + require.True(t, ipv4ChecksumValid(b[:ipv4.HeaderLen])) + require.Equal(t, dir.packet[ipv4.HeaderLen:], b[ipv4.HeaderLen:n]) + } else { + require.Equal(t, dir.packet[ipv6.HeaderLen:], b[ipv6.HeaderLen:n]) + } + } + }) + } + + t.Run("in order both ways at once", func(t *testing.T) { + const count = 2000 + var wg sync.WaitGroup + for _, dir := range []struct { + from, to *Conn + src, dst netip.Addr + }{ + {client, server, clientV4, testDst4}, + {server, client, testDst4, clientV4}, + } { + wg.Go(func() { + for i := range count { + payload := make([]byte, 1200) + payload[0], payload[1] = byte(i>>8), byte(i) + if _, err := dir.from.WritePacket(ipv4Packet(64, 17, dir.src, dir.dst, nil, payload)); !assert.NoError(t, err) { + return + } + } + }) + wg.Go(func() { + b := make([]byte, 1500) + for i := range count { + n, err := dir.to.ReadPacket(b) + if !assert.NoError(t, err) || !assert.Equal(t, ipv4.HeaderLen+1200, n) { + return + } + if !assert.Equal(t, i, int(b[ipv4.HeaderLen])<<8|int(b[ipv4.HeaderLen+1])) { + return + } + } + }) + } + wg.Wait() + }) +} + +func TestHTTP2AddressRequest(t *testing.T) { + client, server := setupHTTP2Conns(t) + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + + _, err := client.RequestAddresses([]netip.Prefix{ + netip.PrefixFrom(netip.IPv4Unspecified(), 32), + netip.PrefixFrom(netip.IPv6Unspecified(), 128), + }) + require.NoError(t, err) + req, err := server.ReceiveAddressRequest(ctx) + require.NoError(t, err) + require.Len(t, req.Prefixes, 2) + require.NoError(t, req.Respond([]netip.Prefix{netip.MustParsePrefix("192.0.2.2/32"), {}}, nil)) + + assigned, err := client.ReceiveAddressAssignment(ctx) + require.NoError(t, err) + require.Len(t, assigned, 2) + require.Equal(t, netip.MustParsePrefix("192.0.2.2/32"), assigned[0].IPPrefix) + require.True(t, assigned[1].Rejected()) +} + +func TestHTTP2Closing(t *testing.T) { + for _, side := range []string{"client", "proxy"} { + t.Run(side, func(t *testing.T) { + client, server := setupHTTP2Conns(t) + closing, peer := client, server + if side == "proxy" { + closing, peer = server, client + } + + require.NoError(t, closing.Close()) + _, err := closing.ReadPacket(make([]byte, 1500)) + require.ErrorIs(t, err, net.ErrClosed) + _, err = closing.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, nil)) + require.ErrorIs(t, err, net.ErrClosed) + + ctx, cancel := context.WithTimeout(t.Context(), 5*time.Second) + defer cancel() + _, err = peer.Routes(ctx) + require.ErrorIs(t, err, net.ErrClosed) + var closeErr *CloseError + require.ErrorAs(t, err, &closeErr) + require.True(t, closeErr.Remote) + _, err = peer.ReadPacket(make([]byte, 1500)) + require.ErrorIs(t, err, net.ErrClosed) + }) + } +} + +func TestHTTP2CloseUnblocksWrites(t *testing.T) { + str, pw := newTestHTTP2Stream() + defer pw.Close() + conn := newProxiedConn(str) + + writeErr := make(chan error, 1) + go func() { + for { + if _, err := conn.WritePacket(ipv4Packet(64, 17, testSrc4, testDst4, nil, make([]byte, 1000))); err != nil { + writeErr <- err + return + } + } + }() + require.Eventually(t, func() bool { + str.body.mu.Lock() + defer str.body.mu.Unlock() + return len(str.body.buf) >= maxStreamBuffer + }, 5*time.Second, time.Millisecond) + + closed := make(chan error, 1) + go func() { closed <- conn.Close() }() + select { + case err := <-closed: + require.NoError(t, err) + case <-time.After(5 * time.Second): + t.Fatal("Close blocked on a stalled stream") + } + require.ErrorIs(t, <-writeErr, net.ErrClosed) +} + +func TestHTTP2DatagramCapsules(t *testing.T) { + str, pw := newTestHTTP2Stream() + defer pw.Close() + conn := newProxiedConn(str) + t.Cleanup(func() { conn.Close() }) + require.NoError(t, conn.AdvertiseRoute([]IPRoute{ + {StartIP: netip.IPv4Unspecified(), EndIP: netip.MustParseAddr("255.255.255.255")}, + })) + + capsule := func(payload []byte) []byte { + b := quicvarint.Append(nil, uint64(capsuleTypeDatagram)) + b = quicvarint.Append(b, uint64(len(payload))) + return append(b, payload...) + } + packet := ipv4Packet(64, 17, testSrc4, testDst4, nil, []byte("foobar")) + go func() { + for _, c := range [][]byte{ + capsule(nil), + capsule([]byte{0x40}), + capsule(append([]byte{0x02}, packet...)), + capsule(append(bytes.Clone(contextIDZero), make([]byte, maxCapsulePacketSize+1)...)), + capsule(append(bytes.Clone(contextIDZero), packet...)), + } { + if _, err := pw.Write(c); err != nil { + return + } + } + }() + b := make([]byte, 1500) + n, err := conn.ReadPacket(b) + require.NoError(t, err) + require.Equal(t, packet, b[:n]) +} + +func TestHTTP2WritesDatagramCapsules(t *testing.T) { + str, pw := newTestHTTP2Stream() + defer pw.Close() + conn := newProxiedConn(str) + t.Cleanup(func() { conn.Close() }) + + packet := ipv4Packet(64, 17, testSrc4, testDst4, nil, []byte("foobar")) + _, err := conn.WritePacket(slices.Clone(packet)) + require.NoError(t, err) + + p := http3.NewCapsuleParser(str.body) + typ, cr, err := p.Next() + require.NoError(t, err) + require.Equal(t, capsuleTypeDatagram, typ) + data, err := io.ReadAll(cr) + require.NoError(t, err) + require.Equal(t, contextIDZero, data[:len(contextIDZero)]) + sent := data[len(contextIDZero):] + require.Len(t, sent, len(packet)) + require.Equal(t, packet[8]-1, sent[8]) + require.Equal(t, packet[ipv4.HeaderLen:], sent[ipv4.HeaderLen:]) +} + +func TestRequestBody(t *testing.T) { + t.Run("coalesces writes", func(t *testing.T) { + b := NewStreamBuffer() + for _, s := range []string{"foo", "bar", "baz"} { + _, err := b.Write([]byte(s)) + require.NoError(t, err) + } + p := make([]byte, 16) + n, err := b.Read(p) + require.NoError(t, err) + require.Equal(t, "foobarbaz", string(p[:n])) + }) + + t.Run("blocks writes while full", func(t *testing.T) { + b := NewStreamBuffer() + _, err := b.Write(make([]byte, maxStreamBuffer)) + require.NoError(t, err) + written := make(chan struct{}) + go func() { + b.Write([]byte("x")) + close(written) + }() + select { + case <-written: + t.Fatal("write did not block") + case <-time.After(50 * time.Millisecond): + } + _, err = b.Read(make([]byte, maxStreamBuffer)) + require.NoError(t, err) + select { + case <-written: + case <-time.After(time.Second): + t.Fatal("write stayed blocked") + } + }) + + t.Run("close", func(t *testing.T) { + b := NewStreamBuffer() + _, err := b.Write([]byte("foo")) + require.NoError(t, err) + require.NoError(t, b.Close()) + _, err = b.Write([]byte("bar")) + require.ErrorIs(t, err, io.ErrClosedPipe) + data, err := io.ReadAll(b) + require.NoError(t, err) + require.Equal(t, "foo", string(data)) + }) + + t.Run("write deadline", func(t *testing.T) { + b := NewStreamBuffer() + _, err := b.Write(make([]byte, maxStreamBuffer)) + require.NoError(t, err) + writeErr := make(chan error, 1) + go func() { + _, err := b.Write([]byte("x")) + writeErr <- err + }() + require.NoError(t, b.SetWriteDeadline(time.Now().Add(50*time.Millisecond))) + select { + case err := <-writeErr: + require.ErrorIs(t, err, os.ErrDeadlineExceeded) + case <-time.After(time.Second): + t.Fatal("write deadline did not unblock the write") + } + require.NoError(t, b.Close()) + data, err := io.ReadAll(b) + require.NoError(t, err) + require.Len(t, data, maxStreamBuffer) + }) + + t.Run("close with error", func(t *testing.T) { + b := NewStreamBuffer() + _, err := b.Write([]byte("foo")) + require.NoError(t, err) + b.CloseWithError(net.ErrClosed) + _, err = b.Read(make([]byte, 16)) + require.ErrorIs(t, err, net.ErrClosed) + _, err = b.Write([]byte("bar")) + require.ErrorIs(t, err, net.ErrClosed) + }) +} + +func TestProxyNeedsAnHTTPStream(t *testing.T) { + _, err := (&Proxy{}).Proxy(httptest.NewRecorder(), &ProxyRequest{}) + require.EqualError(t, err, "connect-ip: response writer is neither an HTTP/3 nor an HTTP/2 stream") +} diff --git a/transport/internet/masque/connectip/proxy.go b/transport/internet/masque/connectip/proxy.go index 323d666df..ab0686ab5 100644 --- a/transport/internet/masque/connectip/proxy.go +++ b/transport/internet/masque/connectip/proxy.go @@ -7,6 +7,7 @@ package connectip import ( + "bufio" "errors" "net/http" @@ -18,13 +19,25 @@ var contextIDZero = quicvarint.Append([]byte{}, 0) type Proxy struct{} -func (s *Proxy) Proxy(w http.ResponseWriter, _ *ProxyRequest) (*Conn, error) { +func (s *Proxy) Proxy(w http.ResponseWriter, r *ProxyRequest) (*Conn, error) { streamer, ok := w.(http3.HTTPStreamer) - if !ok { - return nil, errors.New("connect-ip: response writer is not an HTTP/3 stream") + if !ok && (r == nil || r.body == nil) { + return nil, errors.New("connect-ip: response writer is neither an HTTP/3 nor an HTTP/2 stream") } w.Header().Set(http3.CapsuleProtocolHeader, capsuleProtocolHeaderValue) w.WriteHeader(http.StatusOK) - return newProxiedConn(streamer.HTTPStream()), nil + if ok { + return newProxiedConn(streamer.HTTPStream()), nil + } + controller := http.NewResponseController(w) + if err := controller.Flush(); err != nil { + return nil, err + } + return newProxiedConn(&http2ResponseStream{ + reader: bufio.NewReader(r.body), + body: r.body, + w: w, + controller: controller, + }), nil } diff --git a/transport/internet/masque/connectip/request.go b/transport/internet/masque/connectip/request.go index 4d2c3d6ea..9091caf4e 100644 --- a/transport/internet/masque/connectip/request.go +++ b/transport/internet/masque/connectip/request.go @@ -10,6 +10,7 @@ import ( "context" "errors" "fmt" + "io" "net/http" "strings" @@ -45,7 +46,9 @@ func (r *Request) Header() http.Header { return r.req.Header } func (r *Request) httpRequest() *http.Request { return r.req } -type ProxyRequest struct{} +type ProxyRequest struct { + body io.ReadCloser +} type ProxyRequestParseError struct { HTTPStatus int @@ -62,10 +65,14 @@ func ParseProxyRequest(r *http.Request) (*ProxyRequest, error) { Err: fmt.Errorf("expected CONNECT request, got %s", r.Method), } } - if r.Proto != requestProtocol { + protocol := r.Proto + if r.ProtoMajor == 2 { + protocol = r.Header.Get(":protocol") + } + if protocol != requestProtocol { return nil, &ProxyRequestParseError{ HTTPStatus: http.StatusNotImplemented, - Err: fmt.Errorf("unexpected protocol: %s", r.Proto), + Err: fmt.Errorf("unexpected protocol: %s", protocol), } } capsuleHeaderValues, ok := r.Header[http3.CapsuleProtocolHeader] @@ -82,6 +89,9 @@ func ParseProxyRequest(r *http.Request) (*ProxyRequest, error) { } } + if r.ProtoMajor == 2 { + return &ProxyRequest{body: r.Body}, nil + } return &ProxyRequest{}, nil } diff --git a/transport/internet/masque/connectip/request_test.go b/transport/internet/masque/connectip/request_test.go index 61c8a9c43..223b8a3ec 100644 --- a/transport/internet/masque/connectip/request_test.go +++ b/transport/internet/masque/connectip/request_test.go @@ -71,6 +71,24 @@ func TestProxyRequestParsing(t *testing.T) { require.Equal(t, http.StatusNotImplemented, err.(*ProxyRequestParseError).HTTPStatus) }) + t.Run("HTTP/2", func(t *testing.T) { + req := newRequest("https://localhost:1234/masque/ip") + req.Proto, req.ProtoMajor = "HTTP/2.0", 2 + req.Header.Set(":protocol", requestProtocol) + r, err := ParseProxyRequest(req) + require.NoError(t, err) + require.Equal(t, &ProxyRequest{body: req.Body}, r) + }) + + t.Run("wrong protocol over HTTP/2", func(t *testing.T) { + req := newRequest("https://localhost:1234/masque") + req.Proto, req.ProtoMajor = "HTTP/2.0", 2 + req.Header.Set(":protocol", "websocket") + _, err := ParseProxyRequest(req) + require.EqualError(t, err, "unexpected protocol: websocket") + require.Equal(t, http.StatusNotImplemented, err.(*ProxyRequestParseError).HTTPStatus) + }) + t.Run("wrong request method", func(t *testing.T) { req := newRequest("https://localhost:1234/masque") req.Method = http.MethodHead diff --git a/transport/internet/masque/dialer.go b/transport/internet/masque/dialer.go index 65c75cc81..b59acd578 100644 --- a/transport/internet/masque/dialer.go +++ b/transport/internet/masque/dialer.go @@ -2,9 +2,12 @@ package masque import ( "context" + "net/http" "net/netip" "reflect" "runtime" + "slices" + "strconv" "strings" "time" @@ -16,12 +19,12 @@ import ( "github.com/xtls/xray-core/common/net/cnc" "github.com/xtls/xray-core/common/utils" "github.com/xtls/xray-core/transport/internet" - "github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/hysteria/congestion" "github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr" "github.com/xtls/xray-core/transport/internet/masque/connectip" "github.com/xtls/xray-core/transport/internet/stat" "github.com/xtls/xray-core/transport/internet/tls" + "golang.org/x/net/http2" ) const ( @@ -35,6 +38,9 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me return nil, errors.New("tls config is nil") } config := streamSettings.ProtocolSettings.(*Config) + if usesHTTP2(tlsConfig) { + return dialHTTP2(ctx, dest, streamSettings, tlsConfig, config) + } dest.Network = net.Network_UDP gotlsConfig := tlsConfig.GetTLSConfig(tls.WithDestination(dest)) @@ -73,7 +79,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me if err != nil { return nil, errors.New("failed to dial to dest").Base(err) } - pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn + pktConn = conn.(*net.PacketConnWrapper).PacketConn udpAddr = conn.RemoteAddr() } else { conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) @@ -81,7 +87,7 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me return nil, errors.New("failed to dial to dest").Base(err) } switch c := conn.(type) { - case *internet.PacketConnWrapper: + case *net.PacketConnWrapper: pktConn = c.PacketConn udpAddr = c.RemoteAddr() case *cnc.Connection: @@ -112,7 +118,10 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me return nil, errors.New("unknown congestion control: ", quicParams.Congestion) } - conn, err := establish(ctx, qconn, config, authority(config, gotlsConfig.ServerName, dest.Port)) + cc := (&http3.Transport{EnableDatagrams: true, DisableCompression: true}).NewClientConn(qconn) + conn, err := establish(ctx, connectip.NewClientConn(cc), quicConn{qconn}, func() { + qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeRequestCanceled), "") + }, config, authority(config, gotlsConfig.ServerName, dest.Port)) if err != nil { qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeNoError), "") return nil, err @@ -120,10 +129,58 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me return conn, nil } -func establish(ctx context.Context, qconn *quic.Conn, config *Config, host string) (*Conn, error) { - stop := context.AfterFunc(ctx, func() { - qconn.CloseWithError(quic.ApplicationErrorCode(http3.ErrCodeRequestCanceled), "") - }) +func usesHTTP2(config *tls.Config) bool { + return slices.Contains(config.NextProtocol, http2.NextProtoTLS) && !slices.Contains(config.NextProtocol, http3.NextProtoH3) +} + +func dialHTTP2(ctx context.Context, dest net.Destination, streamSettings *internet.MemoryStreamConfig, tlsConfig *tls.Config, config *Config) (stat.Connection, error) { + dest.Network = net.Network_TCP + gotlsConfig := tlsConfig.GetTLSConfig(tls.WithDestination(dest)) + + var conn net.Conn + var err error + if streamSettings.FinalMask != nil { + conn, err = streamSettings.FinalMask.DialTCP(ctx, dest) + } else { + conn, err = internet.DialSystem(ctx, dest, streamSettings.SocketSettings) + } + if err != nil { + return nil, errors.New("failed to dial to dest").Base(err) + } + if fingerprint := tls.GetFingerprint(tlsConfig.Fingerprint); fingerprint != nil { + conn = tls.UClient(conn, gotlsConfig, fingerprint) + } else { + conn = tls.Client(conn, gotlsConfig) + } + tlsConn := conn.(tls.Interface) + if err := tlsConn.HandshakeContext(ctx); err != nil { + conn.Close() + return nil, err + } + if protocol := tlsConn.NegotiatedProtocol(); protocol != http2.NextProtoTLS { + conn.Close() + return nil, errors.New("the server negotiated ", strconv.Quote(protocol), " instead of h2") + } + + cc, err := newHTTP2ClientConn(conn) + if err != nil { + conn.Close() + return nil, err + } + mconn, err := establish(ctx, connectip.NewHTTP2ClientConn(cc), cc, func() { cc.Close() }, config, authority(config, gotlsConfig.ServerName, dest.Port)) + if err != nil { + cc.Close() + return nil, err + } + return mconn, nil +} + +type tunnelClient interface { + Dial(*connectip.Request) (*connectip.Conn, *http.Response, error) +} + +func establish(ctx context.Context, client tunnelClient, hconn httpConn, abort func(), config *Config, host string) (*Conn, error) { + stop := context.AfterFunc(ctx, abort) defer stop() req, err := connectip.NewRequest(ctx, "https://"+host+config.Path) @@ -151,8 +208,7 @@ func establish(ctx context.Context, qconn *quic.Conn, config *Config, host strin header.Del("User-Agent") } - cc := (&http3.Transport{EnableDatagrams: true, DisableCompression: true}).NewClientConn(qconn) - ipConn, _, err := connectip.NewClientConn(cc).Dial(req) + ipConn, _, err := client.Dial(req) if err != nil { if ctx.Err() != nil { err = context.Cause(ctx) @@ -188,7 +244,7 @@ func establish(ctx context.Context, qconn *quic.Conn, config *Config, host strin conn := &Conn{ ipConn: ipConn, - quicConn: qconn, + httpConn: hconn, local: local, } go conn.serveAddressAssignments() diff --git a/transport/internet/masque/dialer_test.go b/transport/internet/masque/dialer_test.go index 20dd15e18..ddeef00b2 100644 --- a/transport/internet/masque/dialer_test.go +++ b/transport/internet/masque/dialer_test.go @@ -7,8 +7,27 @@ import ( "github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/transport/internet/masque/connectip" + "github.com/xtls/xray-core/transport/internet/tls" ) +func TestUsesHTTP2(t *testing.T) { + for _, c := range []struct { + alpn []string + want bool + }{ + {alpn: nil, want: false}, + {alpn: []string{"h3"}, want: false}, + {alpn: []string{"h2"}, want: true}, + {alpn: []string{"h2", "http/1.1"}, want: true}, + {alpn: []string{"h3", "h2"}, want: false}, + {alpn: []string{"http/1.1"}, want: false}, + } { + if got := usesHTTP2(&tls.Config{NextProtocol: c.alpn}); got != c.want { + t.Errorf("usesHTTP2(%q) = %v, want %v", c.alpn, got, c.want) + } + } +} + func TestAuthority(t *testing.T) { for _, c := range []struct { host, serverName string diff --git a/transport/internet/masque/http2.go b/transport/internet/masque/http2.go new file mode 100644 index 000000000..4665aeaf4 --- /dev/null +++ b/transport/internet/masque/http2.go @@ -0,0 +1,600 @@ +package masque + +import ( + "bufio" + "bytes" + "context" + go_errors "errors" + "io" + "maps" + "net" + "net/http" + "slices" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/xtls/xray-core/common/errors" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +const ( + http2StreamID = 1 + http2DefaultWindow = 65535 + http2DefaultFrameSize = 16 << 10 + http2HeaderTableSize = 64 << 10 + http2StreamWindow = 6 << 20 + http2ConnectionWindow = 15 << 20 + http2MaxHeaderListSize = 256 << 10 + http2WindowUpdateSize = 1 << 20 + http2KeepAlivePeriod = 10 * time.Second + http2IdleTimeout = 30 * time.Second + http2DefaultUserAgent = "Go-http-client/2.0" +) + +var ( + errHTTP2StreamUsed = go_errors.New("http2: the connection carries a single stream") + errHTTP2NoExtendedConnect = go_errors.New("http2: the server did not enable extended CONNECT") + errHTTP2BodyClosed = go_errors.New("http2: response body closed") + errHTTP2IdleTimeout = go_errors.New("http2: no frame received within the idle timeout") +) + +type http2ClientConn struct { + conn net.Conn + + wmu sync.Mutex + bw *bufio.Writer + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder + + lastFrame atomic.Int64 + settings chan struct{} + responses chan *http.Response + aborted chan struct{} + done chan struct{} + + mu sync.Mutex + cond sync.Cond + err error + gotSettings bool + extendedConnect bool + maxFrameSize uint32 + initialWindow int64 + connSendWindow int64 + streamSendWindow int64 + connRecvWindow int64 + streamRecvWindow int64 + streamOpen bool + gotResponse bool + sentEnd bool + recvEnd bool + streamErr error + reqBody io.Closer + recv bytes.Buffer + recvErr error + recvUnacked int64 +} + +func newHTTP2ClientConn(conn net.Conn) (*http2ClientConn, error) { + c := &http2ClientConn{ + conn: conn, + bw: bufio.NewWriter(conn), + settings: make(chan struct{}), + responses: make(chan *http.Response, 1), + aborted: make(chan struct{}), + done: make(chan struct{}), + maxFrameSize: http2DefaultFrameSize, + initialWindow: http2DefaultWindow, + connSendWindow: http2DefaultWindow, + connRecvWindow: http2ConnectionWindow, + streamRecvWindow: http2StreamWindow, + } + c.cond.L = &c.mu + c.fr = http2.NewFramer(c.bw, bufio.NewReader(conn)) + c.fr.SetMaxReadFrameSize(http2DefaultFrameSize) + c.henc = hpack.NewEncoder(&c.hbuf) + c.henc.SetMaxDynamicTableSizeLimit(0) + c.fr.ReadMetaHeaders = hpack.NewDecoder(http2HeaderTableSize, nil) + c.fr.MaxHeaderListSize = http2MaxHeaderListSize + c.lastFrame.Store(time.Now().UnixNano()) + + if err := c.write(func(fr *http2.Framer) error { + if _, err := c.bw.WriteString(http2.ClientPreface); err != nil { + return err + } + if err := fr.WriteSettings( + http2.Setting{ID: http2.SettingHeaderTableSize, Val: http2HeaderTableSize}, + http2.Setting{ID: http2.SettingEnablePush, Val: 0}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: http2StreamWindow}, + http2.Setting{ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize}, + ); err != nil { + return err + } + return fr.WriteWindowUpdate(0, http2ConnectionWindow-http2DefaultWindow) + }); err != nil { + return nil, err + } + go c.readLoop() + go c.keepAlive() + return c, nil +} + +func (c *http2ClientConn) LocalAddr() net.Addr { + return c.conn.LocalAddr() +} + +func (c *http2ClientConn) RemoteAddr() net.Addr { + return c.conn.RemoteAddr() +} + +func (c *http2ClientConn) Close() error { + c.fail(net.ErrClosed) + return nil +} + +func (c *http2ClientConn) RoundTrip(req *http.Request) (*http.Response, error) { + rsp, err := c.roundTrip(req) + if err != nil && req.Body != nil { + req.Body.Close() + } + return rsp, err +} + +func (c *http2ClientConn) roundTrip(req *http.Request) (*http.Response, error) { + ctx := req.Context() + select { + case <-c.settings: + case <-c.done: + return nil, c.connErr() + case <-ctx.Done(): + return nil, context.Cause(ctx) + } + + c.mu.Lock() + switch { + case c.err != nil: + err := c.err + c.mu.Unlock() + return nil, err + case c.streamOpen: + c.mu.Unlock() + return nil, errHTTP2StreamUsed + case req.Header.Get(":protocol") != "" && !c.extendedConnect: + c.mu.Unlock() + return nil, errHTTP2NoExtendedConnect + } + c.streamOpen = true + c.streamSendWindow = c.initialWindow + c.reqBody = req.Body + maxFrameSize := int(c.maxFrameSize) + c.mu.Unlock() + + if err := c.writeHeaders(req, maxFrameSize); err != nil { + c.fail(err) + return nil, err + } + if req.Body != nil { + go c.writeBody(req.Body) + } else { + c.endStream() + } + context.AfterFunc(ctx, func() { c.abortStream(context.Cause(ctx), true) }) + + select { + case rsp := <-c.responses: + return rsp, nil + case <-c.aborted: + c.mu.Lock() + err := c.streamErr + c.mu.Unlock() + return nil, err + } +} + +func (c *http2ClientConn) writeHeaders(req *http.Request, maxFrameSize int) error { + c.wmu.Lock() + defer c.wmu.Unlock() + + c.hbuf.Reset() + field := func(name, value string) { + c.henc.WriteField(hpack.HeaderField{Name: name, Value: value}) + } + host := req.Host + if host == "" { + host = req.URL.Host + } + field(":method", req.Method) + field(":authority", host) + field(":scheme", req.URL.Scheme) + field(":path", req.URL.RequestURI()) + if protocol := req.Header.Get(":protocol"); protocol != "" { + field(":protocol", protocol) + } + if _, ok := req.Header["User-Agent"]; !ok { + field("user-agent", http2DefaultUserAgent) + } + for _, k := range slices.Sorted(maps.Keys(req.Header)) { + name := strings.ToLower(k) + switch name { + case ":protocol", "host", "connection", "proxy-connection", "keep-alive", "transfer-encoding", "upgrade", "content-length": + continue + } + for _, v := range req.Header[k] { + if name == "user-agent" && v == "" { + continue + } + field(name, v) + } + } + + block := c.hbuf.Bytes() + for first := true; first || len(block) > 0; first = false { + chunk := block[:min(len(block), maxFrameSize)] + block = block[len(chunk):] + var err error + if first { + err = c.fr.WriteHeaders(http2.HeadersFrameParam{StreamID: http2StreamID, BlockFragment: chunk, EndHeaders: len(block) == 0}) + } else { + err = c.fr.WriteContinuation(http2StreamID, len(block) == 0, chunk) + } + if err != nil { + return err + } + } + return c.bw.Flush() +} + +func (c *http2ClientConn) writeBody(body io.ReadCloser) { + defer body.Close() + buf := make([]byte, http2DefaultFrameSize) + for { + n, err := body.Read(buf) + for data := buf[:n]; len(data) > 0; { + allowed, err := c.awaitSendWindow(len(data)) + if err != nil { + return + } + if err := c.write(func(fr *http2.Framer) error { + return fr.WriteData(http2StreamID, false, data[:allowed]) + }); err != nil { + c.fail(err) + return + } + data = data[allowed:] + } + if err == io.EOF { + c.endStream() + return + } + if err != nil { + c.abortStream(err, true) + return + } + } +} + +func (c *http2ClientConn) awaitSendWindow(n int) (int, error) { + c.mu.Lock() + defer c.mu.Unlock() + for { + if c.streamErr != nil { + return 0, c.streamErr + } + if window := min(c.connSendWindow, c.streamSendWindow); window > 0 { + n = int(min(int64(n), window, int64(c.maxFrameSize))) + c.connSendWindow -= int64(n) + c.streamSendWindow -= int64(n) + return n, nil + } + c.cond.Wait() + } +} + +func (c *http2ClientConn) endStream() { + c.mu.Lock() + if c.streamErr != nil || c.sentEnd { + c.mu.Unlock() + return + } + c.sentEnd = true + c.mu.Unlock() + if err := c.write(func(fr *http2.Framer) error { + return fr.WriteData(http2StreamID, true, nil) + }); err != nil { + c.fail(err) + } +} + +func (c *http2ClientConn) write(f func(*http2.Framer) error) error { + c.wmu.Lock() + defer c.wmu.Unlock() + if err := f(c.fr); err != nil { + return err + } + return c.bw.Flush() +} + +func (c *http2ClientConn) connErr() error { + c.mu.Lock() + defer c.mu.Unlock() + return c.err +} + +func (c *http2ClientConn) fail(err error) { + c.mu.Lock() + if c.err == nil { + c.err = err + } + c.mu.Unlock() + c.abortStream(err, false) + c.conn.Close() +} + +func (c *http2ClientConn) abortStream(err error, reset bool) { + c.mu.Lock() + if c.streamErr != nil { + c.mu.Unlock() + return + } + c.streamErr = err + if c.recvErr == nil { + c.recvErr = err + } + reset = reset && c.streamOpen && !(c.sentEnd && c.recvEnd) + body := c.reqBody + close(c.aborted) + c.cond.Broadcast() + c.mu.Unlock() + + if body != nil { + body.Close() + } + if reset { + go c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(http2StreamID, http2.ErrCodeCancel) + }) + } +} + +func (c *http2ClientConn) keepAlive() { + ticker := time.NewTicker(http2KeepAlivePeriod) + defer ticker.Stop() + for { + select { + case <-c.done: + return + case <-ticker.C: + } + idle := time.Since(time.Unix(0, c.lastFrame.Load())) + if idle >= http2IdleTimeout { + c.fail(errHTTP2IdleTimeout) + return + } + if idle >= http2KeepAlivePeriod { + go c.write(func(fr *http2.Framer) error { + return fr.WritePing(false, [8]byte{}) + }) + } + } +} + +func (c *http2ClientConn) readLoop() { + defer close(c.done) + for { + f, err := c.fr.ReadFrame() + if err != nil { + var streamErr http2.StreamError + if go_errors.As(err, &streamErr) && streamErr.StreamID == http2StreamID { + c.abortStream(streamErr, true) + continue + } + c.fail(err) + return + } + c.lastFrame.Store(time.Now().UnixNano()) + if err := c.handleFrame(f); err != nil { + c.fail(err) + return + } + } +} + +func (c *http2ClientConn) handleFrame(f http2.Frame) error { + switch f := f.(type) { + case *http2.SettingsFrame: + if f.IsAck() { + return nil + } + if err := c.applySettings(f); err != nil { + return err + } + return c.write((*http2.Framer).WriteSettingsAck) + case *http2.PingFrame: + if f.IsAck() { + return nil + } + return c.write(func(fr *http2.Framer) error { + return fr.WritePing(true, f.Data) + }) + case *http2.WindowUpdateFrame: + c.mu.Lock() + switch f.StreamID { + case 0: + c.connSendWindow += int64(f.Increment) + case http2StreamID: + c.streamSendWindow += int64(f.Increment) + } + c.cond.Broadcast() + c.mu.Unlock() + case *http2.MetaHeadersFrame: + if f.StreamID == http2StreamID { + c.handleHeaders(f) + } + case *http2.DataFrame: + return c.handleData(f) + case *http2.RSTStreamFrame: + if f.StreamID == http2StreamID { + c.abortStream(http2.StreamError{StreamID: f.StreamID, Code: f.ErrCode}, false) + } + case *http2.GoAwayFrame: + if f.ErrCode != http2.ErrCodeNo || f.LastStreamID < http2StreamID { + return errors.New("http2: the server sent GOAWAY (", f.ErrCode, ")") + } + case *http2.PushPromiseFrame: + return http2.ConnectionError(http2.ErrCodeProtocol) + } + return nil +} + +func (c *http2ClientConn) applySettings(f *http2.SettingsFrame) error { + c.mu.Lock() + defer c.mu.Unlock() + if err := f.ForeachSetting(func(s http2.Setting) error { + if err := s.Valid(); err != nil { + return err + } + switch s.ID { + case http2.SettingMaxFrameSize: + c.maxFrameSize = s.Val + case http2.SettingInitialWindowSize: + c.streamSendWindow += int64(s.Val) - c.initialWindow + c.initialWindow = int64(s.Val) + case http2.SettingEnableConnectProtocol: + if !c.gotSettings { + c.extendedConnect = s.Val == 1 + } + } + return nil + }); err != nil { + return err + } + if !c.gotSettings { + c.gotSettings = true + close(c.settings) + } + c.cond.Broadcast() + return nil +} + +func (c *http2ClientConn) handleHeaders(f *http2.MetaHeadersFrame) { + c.mu.Lock() + gotResponse := c.gotResponse + c.mu.Unlock() + if !gotResponse { + status, err := strconv.Atoi(f.PseudoValue("status")) + if err != nil || status < 100 || status > 999 { + c.abortStream(errors.New("http2: invalid response status ", strconv.Quote(f.PseudoValue("status"))), true) + return + } + if status < 200 { + return + } + header := make(http.Header) + for _, hf := range f.RegularFields() { + header.Add(hf.Name, hf.Value) + } + c.mu.Lock() + c.gotResponse = true + c.mu.Unlock() + c.responses <- &http.Response{ + Status: strconv.Itoa(status) + " " + http.StatusText(status), + StatusCode: status, + Proto: "HTTP/2.0", + ProtoMajor: 2, + Header: header, + Body: &http2ResponseBody{c}, + ContentLength: -1, + } + } + if f.StreamEnded() { + c.mu.Lock() + c.recvEnd = true + if c.recvErr == nil { + c.recvErr = io.EOF + } + c.cond.Broadcast() + c.mu.Unlock() + } +} + +func (c *http2ClientConn) handleData(f *http2.DataFrame) error { + size := int64(f.Length) + c.mu.Lock() + c.connRecvWindow -= size + if c.connRecvWindow < 0 { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + if f.StreamID != http2StreamID || c.recvErr != nil { + c.connRecvWindow += size + c.mu.Unlock() + if size == 0 { + return nil + } + return c.write(func(fr *http2.Framer) error { + return fr.WriteWindowUpdate(0, uint32(size)) + }) + } + c.streamRecvWindow -= size + if c.streamRecvWindow < 0 { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + c.recv.Write(f.Data()) + c.recvUnacked += size - int64(len(f.Data())) + if f.StreamEnded() { + c.recvEnd = true + c.recvErr = io.EOF + } + c.cond.Broadcast() + c.mu.Unlock() + return nil +} + +type http2ResponseBody struct { + c *http2ClientConn +} + +func (b *http2ResponseBody) Read(p []byte) (int, error) { + c := b.c + c.mu.Lock() + for c.recv.Len() == 0 && c.recvErr == nil { + c.cond.Wait() + } + if c.recv.Len() == 0 { + err := c.recvErr + c.mu.Unlock() + return 0, err + } + n, _ := c.recv.Read(p) + c.recvUnacked += int64(n) + var update int64 + if c.recvUnacked >= http2WindowUpdateSize && !c.recvEnd { + update = c.recvUnacked + c.recvUnacked = 0 + c.connRecvWindow += update + c.streamRecvWindow += update + } + c.mu.Unlock() + + if update > 0 { + if err := c.write(func(fr *http2.Framer) error { + if err := fr.WriteWindowUpdate(0, uint32(update)); err != nil { + return err + } + return fr.WriteWindowUpdate(http2StreamID, uint32(update)) + }); err != nil { + c.fail(err) + } + } + return n, nil +} + +func (b *http2ResponseBody) Close() error { + b.c.abortStream(errHTTP2BodyClosed, true) + return nil +} diff --git a/transport/internet/masque/http2_server.go b/transport/internet/masque/http2_server.go new file mode 100644 index 000000000..98aaa25bf --- /dev/null +++ b/transport/internet/masque/http2_server.go @@ -0,0 +1,680 @@ +package masque + +import ( + "bufio" + "bytes" + "context" + go_errors "errors" + "io" + "maps" + "math" + "net" + "net/http" + "net/url" + "slices" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/xtls/xray-core/transport/internet/masque/connectip" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +const ( + http2MaxConcurrentStreams = 100 + http2HandshakeTimeout = 10 * time.Second + http2DefaultHeaderTable = 4096 +) + +var ( + errHTTP2BadPreface = go_errors.New("http2: invalid connection preface") + errHTTP2BadRequest = go_errors.New("http2: malformed request") + errHTTP2StreamClosed = go_errors.New("http2: stream closed") + errHTTP2RequestBodyClosed = go_errors.New("http2: request body closed") +) + +type connAddrsKey struct{} + +type connAddrs struct { + local net.Addr + remote net.Addr +} + +type http2ServerConn struct { + conn net.Conn + handler http.Handler + ctx context.Context + cancel context.CancelFunc + + wmu sync.Mutex + bw *bufio.Writer + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder + + lastFrame atomic.Int64 + + mu sync.Mutex + cond sync.Cond + err error + maxFrameSize uint32 + initialWindow int64 + connSendWindow int64 + connRecvWindow int64 + connRecvUnacked int64 + streams map[uint32]*http2ServerStream + lastStreamID uint32 +} + +type http2ServerStream struct { + c *http2ServerConn + id uint32 + ctx context.Context + cancel context.CancelFunc + header http.Header + out *connectip.StreamBuffer + sent chan struct{} + + sendWindow int64 + recvWindow int64 + recvUnacked int64 + recv bytes.Buffer + recvEnd bool + wroteHeader bool + sentEnd bool + resetErr error +} + +func serveHTTP2(ctx context.Context, conn net.Conn, handler http.Handler) { + ctx, cancel := context.WithCancel(ctx) + c := &http2ServerConn{ + conn: conn, + handler: handler, + ctx: context.WithValue(ctx, connAddrsKey{}, connAddrs{local: conn.LocalAddr(), remote: conn.RemoteAddr()}), + cancel: cancel, + bw: bufio.NewWriter(conn), + maxFrameSize: http2DefaultFrameSize, + initialWindow: http2DefaultWindow, + connSendWindow: http2DefaultWindow, + connRecvWindow: http2ConnectionWindow, + streams: make(map[uint32]*http2ServerStream), + } + c.cond.L = &c.mu + c.henc = hpack.NewEncoder(&c.hbuf) + c.henc.SetMaxDynamicTableSizeLimit(0) + c.lastFrame.Store(time.Now().UnixNano()) + + br := bufio.NewReader(conn) + preface := make([]byte, len(http2.ClientPreface)) + conn.SetReadDeadline(time.Now().Add(http2HandshakeTimeout)) + if _, err := io.ReadFull(br, preface); err != nil || string(preface) != http2.ClientPreface { + c.fail(errHTTP2BadPreface) + return + } + conn.SetReadDeadline(time.Time{}) + + c.fr = http2.NewFramer(c.bw, br) + c.fr.SetMaxReadFrameSize(http2DefaultFrameSize) + c.fr.ReadMetaHeaders = hpack.NewDecoder(http2DefaultHeaderTable, nil) + c.fr.MaxHeaderListSize = http2MaxHeaderListSize + if err := c.write(func(fr *http2.Framer) error { + if err := fr.WriteSettings( + http2.Setting{ID: http2.SettingMaxConcurrentStreams, Val: http2MaxConcurrentStreams}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: http2StreamWindow}, + http2.Setting{ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize}, + http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}, + ); err != nil { + return err + } + return fr.WriteWindowUpdate(0, http2ConnectionWindow-http2DefaultWindow) + }); err != nil { + c.fail(err) + return + } + go c.keepAlive() + c.readLoop() +} + +func (c *http2ServerConn) write(f func(*http2.Framer) error) error { + c.wmu.Lock() + defer c.wmu.Unlock() + if err := f(c.fr); err != nil { + return err + } + return c.bw.Flush() +} + +func (c *http2ServerConn) fail(err error) { + c.mu.Lock() + if c.err == nil { + c.err = err + } + streams := slices.Collect(maps.Values(c.streams)) + c.cond.Broadcast() + c.mu.Unlock() + for _, st := range streams { + st.out.CloseWithError(err) + st.cancel() + } + c.cancel() + c.conn.Close() +} + +func (c *http2ServerConn) keepAlive() { + ticker := time.NewTicker(http2KeepAlivePeriod) + defer ticker.Stop() + for { + select { + case <-c.ctx.Done(): + return + case <-ticker.C: + } + idle := time.Since(time.Unix(0, c.lastFrame.Load())) + if idle >= http2IdleTimeout { + c.fail(errHTTP2IdleTimeout) + return + } + if idle >= http2KeepAlivePeriod { + go c.write(func(fr *http2.Framer) error { + return fr.WritePing(false, [8]byte{}) + }) + } + } +} + +func (c *http2ServerConn) readLoop() { + for { + f, err := c.fr.ReadFrame() + if err != nil { + var streamErr http2.StreamError + if go_errors.As(err, &streamErr) { + if err := c.resetStream(streamErr); err != nil { + c.fail(err) + return + } + continue + } + c.fail(err) + return + } + c.lastFrame.Store(time.Now().UnixNano()) + if err := c.handleFrame(f); err != nil { + c.fail(err) + return + } + } +} + +func (c *http2ServerConn) handleFrame(f http2.Frame) error { + switch f := f.(type) { + case *http2.SettingsFrame: + if f.IsAck() { + return nil + } + if err := c.applySettings(f); err != nil { + return err + } + return c.write((*http2.Framer).WriteSettingsAck) + case *http2.PingFrame: + if f.IsAck() { + return nil + } + return c.write(func(fr *http2.Framer) error { + return fr.WritePing(true, f.Data) + }) + case *http2.WindowUpdateFrame: + return c.handleWindowUpdate(f) + case *http2.MetaHeadersFrame: + return c.handleHeaders(f) + case *http2.DataFrame: + return c.handleData(f) + case *http2.RSTStreamFrame: + c.abortStream(f.StreamID, http2.StreamError{StreamID: f.StreamID, Code: f.ErrCode}, false) + case *http2.PushPromiseFrame: + return http2.ConnectionError(http2.ErrCodeProtocol) + } + return nil +} + +func (c *http2ServerConn) applySettings(f *http2.SettingsFrame) error { + c.mu.Lock() + defer c.mu.Unlock() + defer c.cond.Broadcast() + return f.ForeachSetting(func(s http2.Setting) error { + if err := s.Valid(); err != nil { + return err + } + switch s.ID { + case http2.SettingMaxFrameSize: + c.maxFrameSize = s.Val + case http2.SettingInitialWindowSize: + delta := int64(s.Val) - c.initialWindow + c.initialWindow = int64(s.Val) + for _, st := range c.streams { + st.sendWindow += delta + if st.sendWindow > math.MaxInt32 { + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + } + } + return nil + }) +} + +func (c *http2ServerConn) handleWindowUpdate(f *http2.WindowUpdateFrame) error { + c.mu.Lock() + defer c.mu.Unlock() + if f.StreamID == 0 { + c.connSendWindow += int64(f.Increment) + if c.connSendWindow > math.MaxInt32 { + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + } else if st := c.streams[f.StreamID]; st != nil { + st.sendWindow += int64(f.Increment) + if st.sendWindow > math.MaxInt32 { + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + } + c.cond.Broadcast() + return nil +} + +func (c *http2ServerConn) handleHeaders(f *http2.MetaHeadersFrame) error { + c.mu.Lock() + if st := c.streams[f.StreamID]; st != nil { + if !f.StreamEnded() { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeProtocol) + } + st.recvEnd = true + c.cond.Broadcast() + c.mu.Unlock() + return nil + } + if f.StreamID%2 == 0 || f.StreamID <= c.lastStreamID { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeProtocol) + } + c.lastStreamID = f.StreamID + refused := len(c.streams) >= http2MaxConcurrentStreams + c.mu.Unlock() + + if refused { + return c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(f.StreamID, http2.ErrCodeRefusedStream) + }) + } + req, err := newHTTP2Request(f) + if err != nil { + return c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(f.StreamID, http2.ErrCodeProtocol) + }) + } + + ctx, cancel := context.WithCancel(c.ctx) + st := &http2ServerStream{ + c: c, + id: f.StreamID, + ctx: ctx, + cancel: cancel, + header: make(http.Header), + out: connectip.NewStreamBuffer(), + sent: make(chan struct{}), + recvWindow: http2StreamWindow, + recvEnd: f.StreamEnded(), + } + req = req.WithContext(ctx) + req.RemoteAddr = c.conn.RemoteAddr().String() + if st.recvEnd { + req.Body = http.NoBody + req.ContentLength = 0 + } else { + req.Body = &http2RequestBody{st: st} + req.ContentLength = -1 + } + + c.mu.Lock() + if c.err != nil { + err := c.err + c.mu.Unlock() + cancel() + return err + } + st.sendWindow = c.initialWindow + c.streams[f.StreamID] = st + c.mu.Unlock() + + go c.serveStream(st, req) + return nil +} + +func newHTTP2Request(f *http2.MetaHeadersFrame) (*http.Request, error) { + method := f.PseudoValue("method") + scheme := f.PseudoValue("scheme") + authority := f.PseudoValue("authority") + path := f.PseudoValue("path") + protocol := f.PseudoValue("protocol") + if method == "" || (protocol != "" && method != http.MethodConnect) { + return nil, errHTTP2BadRequest + } + var u *url.URL + if method == http.MethodConnect && protocol == "" { + if authority == "" || path != "" || scheme != "" { + return nil, errHTTP2BadRequest + } + u = &url.URL{Host: authority} + } else { + if path == "" || scheme == "" { + return nil, errHTTP2BadRequest + } + var err error + if u, err = url.ParseRequestURI(path); err != nil { + return nil, errHTTP2BadRequest + } + } + header := make(http.Header) + for _, hf := range f.RegularFields() { + header.Add(hf.Name, hf.Value) + } + if protocol != "" { + header.Set(":protocol", protocol) + } + if authority == "" { + authority = header.Get("Host") + } + return &http.Request{ + Method: method, + URL: u, + Proto: "HTTP/2.0", + ProtoMajor: 2, + Header: header, + Host: authority, + RequestURI: path, + }, nil +} + +func (c *http2ServerConn) handleData(f *http2.DataFrame) error { + size := int64(f.Length) + c.mu.Lock() + c.connRecvWindow -= size + if c.connRecvWindow < 0 { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + st := c.streams[f.StreamID] + if st == nil || st.resetErr != nil || st.recvEnd { + idle := st == nil && f.StreamID > c.lastStreamID + c.connRecvWindow += size + c.mu.Unlock() + if idle { + return http2.ConnectionError(http2.ErrCodeProtocol) + } + if size == 0 { + return nil + } + return c.write(func(fr *http2.Framer) error { + return fr.WriteWindowUpdate(0, uint32(size)) + }) + } + st.recvWindow -= size + if st.recvWindow < 0 { + c.mu.Unlock() + return http2.ConnectionError(http2.ErrCodeFlowControl) + } + st.recv.Write(f.Data()) + padding := size - int64(len(f.Data())) + st.recvUnacked += padding + c.connRecvUnacked += padding + if f.StreamEnded() { + st.recvEnd = true + } + c.cond.Broadcast() + c.mu.Unlock() + return nil +} + +func (c *http2ServerConn) resetStream(streamErr http2.StreamError) error { + c.mu.Lock() + if streamErr.StreamID%2 == 1 && streamErr.StreamID > c.lastStreamID { + c.lastStreamID = streamErr.StreamID + } + c.mu.Unlock() + c.abortStream(streamErr.StreamID, streamErr, false) + return c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(streamErr.StreamID, streamErr.Code) + }) +} + +func (c *http2ServerConn) abortStream(id uint32, err error, reset bool) { + c.mu.Lock() + st := c.streams[id] + if st == nil || st.resetErr != nil { + c.mu.Unlock() + return + } + st.resetErr = err + c.cond.Broadcast() + c.mu.Unlock() + st.out.CloseWithError(err) + st.cancel() + if reset { + go c.write(func(fr *http2.Framer) error { + return fr.WriteRSTStream(id, http2.ErrCodeCancel) + }) + } +} + +func (c *http2ServerConn) serveStream(st *http2ServerStream, req *http.Request) { + go st.sendLoop() + c.handler.ServeHTTP(&http2ResponseWriter{st: st}, req) + st.writeHeader(http.StatusOK) + st.out.Close() + <-st.sent + + c.mu.Lock() + finish := st.resetErr == nil && !st.sentEnd && c.err == nil + refuse := finish && !st.recvEnd + st.sentEnd = true + delete(c.streams, st.id) + c.cond.Broadcast() + c.mu.Unlock() + st.cancel() + + if finish { + if err := c.write(func(fr *http2.Framer) error { + if err := fr.WriteData(st.id, true, nil); err != nil { + return err + } + if refuse { + return fr.WriteRSTStream(st.id, http2.ErrCodeNo) + } + return nil + }); err != nil { + c.fail(err) + } + } +} + +func (st *http2ServerStream) writeHeader(code int) { + c := st.c + c.mu.Lock() + if st.wroteHeader || st.resetErr != nil || c.err != nil { + c.mu.Unlock() + return + } + st.wroteHeader = true + header := st.header.Clone() + maxFrameSize := int(c.maxFrameSize) + c.mu.Unlock() + + if err := c.write(func(fr *http2.Framer) error { + c.hbuf.Reset() + c.henc.WriteField(hpack.HeaderField{Name: ":status", Value: strconv.Itoa(code)}) + for _, k := range slices.Sorted(maps.Keys(header)) { + name := strings.ToLower(k) + switch name { + case "connection", "proxy-connection", "keep-alive", "transfer-encoding", "upgrade": + continue + } + for _, v := range header[k] { + c.henc.WriteField(hpack.HeaderField{Name: name, Value: v}) + } + } + block := c.hbuf.Bytes() + for first := true; first || len(block) > 0; first = false { + chunk := block[:min(len(block), maxFrameSize)] + block = block[len(chunk):] + var err error + if first { + err = fr.WriteHeaders(http2.HeadersFrameParam{StreamID: st.id, BlockFragment: chunk, EndHeaders: len(block) == 0}) + } else { + err = fr.WriteContinuation(st.id, len(block) == 0, chunk) + } + if err != nil { + return err + } + } + return nil + }); err != nil { + c.fail(err) + } +} + +func (st *http2ServerStream) sendLoop() { + defer close(st.sent) + c := st.c + buf := make([]byte, http2DefaultFrameSize) + for { + n, err := st.out.Read(buf) + for data := buf[:n]; len(data) > 0; { + allowed, err := st.awaitSendWindow(len(data)) + if err != nil { + st.out.CloseWithError(err) + return + } + if err := c.write(func(fr *http2.Framer) error { + return fr.WriteData(st.id, false, data[:allowed]) + }); err != nil { + c.fail(err) + return + } + data = data[allowed:] + } + if err != nil { + return + } + } +} + +func (st *http2ServerStream) awaitSendWindow(n int) (int, error) { + c := st.c + c.mu.Lock() + defer c.mu.Unlock() + for { + switch { + case st.resetErr != nil: + return 0, st.resetErr + case c.err != nil: + return 0, c.err + case st.sentEnd: + return 0, errHTTP2StreamClosed + } + if window := min(c.connSendWindow, st.sendWindow); window > 0 { + n = int(min(int64(n), window, int64(c.maxFrameSize))) + c.connSendWindow -= int64(n) + st.sendWindow -= int64(n) + return n, nil + } + c.cond.Wait() + } +} + +type http2ResponseWriter struct { + st *http2ServerStream +} + +func (w *http2ResponseWriter) Header() http.Header { return w.st.header } + +func (w *http2ResponseWriter) WriteHeader(code int) { w.st.writeHeader(code) } + +func (w *http2ResponseWriter) Write(p []byte) (int, error) { + w.st.writeHeader(http.StatusOK) + return w.st.out.Write(p) +} + +func (w *http2ResponseWriter) Flush() { w.st.writeHeader(http.StatusOK) } + +func (w *http2ResponseWriter) SetWriteDeadline(t time.Time) error { + return w.st.out.SetWriteDeadline(t) +} + +type http2RequestBody struct { + st *http2ServerStream +} + +func (b *http2RequestBody) Read(p []byte) (int, error) { + st := b.st + c := st.c + c.mu.Lock() + for st.recv.Len() == 0 && !st.recvEnd && st.resetErr == nil && c.err == nil { + c.cond.Wait() + } + if st.recv.Len() == 0 { + err := io.EOF + switch { + case st.resetErr != nil: + err = st.resetErr + case c.err != nil && !st.recvEnd: + err = c.err + } + c.mu.Unlock() + return 0, err + } + n, _ := st.recv.Read(p) + st.recvUnacked += int64(n) + c.connRecvUnacked += int64(n) + var streamUpdate, connUpdate int64 + if st.recvUnacked >= http2WindowUpdateSize && !st.recvEnd { + streamUpdate = st.recvUnacked + st.recvUnacked = 0 + st.recvWindow += streamUpdate + } + if c.connRecvUnacked >= http2WindowUpdateSize { + connUpdate = c.connRecvUnacked + c.connRecvUnacked = 0 + c.connRecvWindow += connUpdate + } + c.mu.Unlock() + + if streamUpdate > 0 || connUpdate > 0 { + if err := c.write(func(fr *http2.Framer) error { + if connUpdate > 0 { + if err := fr.WriteWindowUpdate(0, uint32(connUpdate)); err != nil { + return err + } + } + if streamUpdate > 0 { + return fr.WriteWindowUpdate(st.id, uint32(streamUpdate)) + } + return nil + }); err != nil { + c.fail(err) + } + } + return n, nil +} + +func (b *http2RequestBody) Close() error { + st := b.st + c := st.c + c.mu.Lock() + done := st.recvEnd + c.mu.Unlock() + if !done { + c.abortStream(st.id, errHTTP2RequestBodyClosed, true) + } + return nil +} diff --git a/transport/internet/masque/http2_server_test.go b/transport/internet/masque/http2_server_test.go new file mode 100644 index 000000000..510e37438 --- /dev/null +++ b/transport/internet/masque/http2_server_test.go @@ -0,0 +1,281 @@ +package masque + +import ( + "bytes" + "context" + "crypto/rand" + "crypto/sha256" + "io" + "net" + "net/http" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +func serveHTTP2Pipe(t *testing.T, handler http.Handler) net.Conn { + t.Helper() + client, server := tcpPipe(t) + done := make(chan struct{}) + go func() { + serveHTTP2(context.Background(), server, handler) + close(done) + }() + t.Cleanup(func() { + client.Close() + server.Close() + <-done + }) + return client +} + +type http2ClientPeer struct { + t *testing.T + conn net.Conn + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder +} + +func newHTTP2ClientPeer(t *testing.T, handler http.Handler) (*http2ClientPeer, []http2.Setting) { + t.Helper() + conn := serveHTTP2Pipe(t, handler) + p := &http2ClientPeer{t: t, conn: conn, fr: http2.NewFramer(conn, conn)} + p.henc = hpack.NewEncoder(&p.hbuf) + p.fr.ReadMetaHeaders = hpack.NewDecoder(4096, nil) + _, err := io.WriteString(conn, http2.ClientPreface) + require.NoError(t, err) + require.NoError(t, p.fr.WriteSettings()) + + f := p.readFrame() + require.IsType(t, &http2.SettingsFrame{}, f) + var settings []http2.Setting + f.(*http2.SettingsFrame).ForeachSetting(func(s http2.Setting) error { + settings = append(settings, s) + return nil + }) + f = p.readFrame() + require.IsType(t, &http2.WindowUpdateFrame{}, f) + require.Equal(t, uint32(http2ConnectionWindow-http2DefaultWindow), f.(*http2.WindowUpdateFrame).Increment) + f = p.readFrame() + require.True(t, f.(*http2.SettingsFrame).IsAck()) + require.NoError(t, p.fr.WriteSettingsAck()) + return p, settings +} + +func (p *http2ClientPeer) readFrame() http2.Frame { + p.t.Helper() + p.conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + f, err := p.fr.ReadFrame() + require.NoError(p.t, err) + return f +} + +func (p *http2ClientPeer) writeHeaders(streamID uint32, endStream bool, fields ...string) { + p.t.Helper() + p.hbuf.Reset() + for i := 0; i < len(fields); i += 2 { + require.NoError(p.t, p.henc.WriteField(hpack.HeaderField{Name: fields[i], Value: fields[i+1]})) + } + require.NoError(p.t, p.fr.WriteHeaders(http2.HeadersFrameParam{ + StreamID: streamID, + BlockFragment: p.hbuf.Bytes(), + EndHeaders: true, + EndStream: endStream, + })) +} + +func (p *http2ClientPeer) writeConnect(streamID uint32) { + p.t.Helper() + p.writeHeaders(streamID, false, + ":method", "CONNECT", + ":protocol", "connect-ip", + ":scheme", "https", + ":authority", "proxy.example", + ":path", "/.well-known/masque/ip/*/*/", + "capsule-protocol", "?1", + ) +} + +func TestHTTP2ServerSettings(t *testing.T) { + _, settings := newHTTP2ClientPeer(t, http.NotFoundHandler()) + require.Equal(t, []http2.Setting{ + {ID: http2.SettingMaxConcurrentStreams, Val: http2MaxConcurrentStreams}, + {ID: http2.SettingInitialWindowSize, Val: http2StreamWindow}, + {ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize}, + {ID: http2.SettingEnableConnectProtocol, Val: 1}, + }, settings) +} + +func TestHTTP2ServerRoundTrip(t *testing.T) { + handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + assert.Equal(t, http.MethodConnect, r.Method) + assert.Equal(t, "connect-ip", r.Header.Get(":protocol")) + assert.Equal(t, "?1", r.Header.Get("Capsule-Protocol")) + assert.Equal(t, "Basic dTpw", r.Header.Get("Authorization")) + assert.Equal(t, 2, r.ProtoMajor) + assert.Equal(t, "proxy.example", r.Host) + assert.Equal(t, "/.well-known/masque/ip/*/*/", r.URL.Path) + w.Header().Set("Capsule-Protocol", "?1") + w.WriteHeader(http.StatusOK) + assert.NoError(t, http.NewResponseController(w).Flush()) + _, err := io.Copy(w, r.Body) + assert.NoError(t, err) + }) + cc, err := newHTTP2ClientConn(serveHTTP2Pipe(t, handler)) + require.NoError(t, err) + defer cc.Close() + + pr, pw := io.Pipe() + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), pr)) + require.NoError(t, err) + require.Equal(t, http.StatusOK, rsp.StatusCode) + require.Equal(t, "?1", rsp.Header.Get("Capsule-Protocol")) + + payload := make([]byte, 3*http2ConnectionWindow/2) + rand.Read(payload) + go func() { + pw.Write(payload) + pw.Close() + }() + echoed := sha256.New() + n, err := io.Copy(echoed, rsp.Body) + require.NoError(t, err) + require.Equal(t, int64(len(payload)), n) + require.Equal(t, sha256.Sum256(payload), [32]byte(echoed.Sum(nil))) +} + +func TestHTTP2ServerStatus(t *testing.T) { + p, _ := newHTTP2ClientPeer(t, http.NotFoundHandler()) + p.writeConnect(1) + f := p.readFrame() + require.IsType(t, &http2.MetaHeadersFrame{}, f) + require.Equal(t, "404", f.(*http2.MetaHeadersFrame).PseudoValue("status")) + var body []byte + for { + f = p.readFrame() + require.IsType(t, &http2.DataFrame{}, f) + body = append(body, f.(*http2.DataFrame).Data()...) + if f.(*http2.DataFrame).StreamEnded() { + break + } + } + require.Equal(t, "404 page not found\n", string(body)) + f = p.readFrame() + require.IsType(t, &http2.RSTStreamFrame{}, f) + require.Equal(t, http2.ErrCodeNo, f.(*http2.RSTStreamFrame).ErrCode) +} + +func TestHTTP2ServerMalformedRequests(t *testing.T) { + for _, tc := range []struct { + name string + fields []string + }{ + {"no method", []string{":scheme", "https", ":path", "/", ":authority", "proxy.example"}}, + {"no path", []string{":method", "GET", ":scheme", "https", ":authority", "proxy.example"}}, + {"protocol without CONNECT", []string{":method", "GET", ":protocol", "connect-ip", ":scheme", "https", ":path", "/", ":authority", "proxy.example"}}, + {"plain CONNECT with a path", []string{":method", "CONNECT", ":path", "/", ":authority", "proxy.example"}}, + } { + t.Run(tc.name, func(t *testing.T) { + p, _ := newHTTP2ClientPeer(t, http.HandlerFunc(func(http.ResponseWriter, *http.Request) { + t.Error("the handler saw a malformed request") + })) + p.writeHeaders(1, false, tc.fields...) + f := p.readFrame() + require.IsType(t, &http2.RSTStreamFrame{}, f) + require.Equal(t, http2.ErrCodeProtocol, f.(*http2.RSTStreamFrame).ErrCode) + }) + } +} + +func TestHTTP2ServerRefusesExtraStreams(t *testing.T) { + release := make(chan struct{}) + var started sync.WaitGroup + started.Add(http2MaxConcurrentStreams) + p, _ := newHTTP2ClientPeer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + started.Done() + <-release + })) + defer close(release) + for i := range http2MaxConcurrentStreams { + p.writeConnect(uint32(2*i + 1)) + } + started.Wait() + p.writeConnect(2*http2MaxConcurrentStreams + 1) + f := p.readFrame() + require.IsType(t, &http2.RSTStreamFrame{}, f) + require.Equal(t, uint32(2*http2MaxConcurrentStreams+1), f.Header().StreamID) + require.Equal(t, http2.ErrCodeRefusedStream, f.(*http2.RSTStreamFrame).ErrCode) +} + +func TestHTTP2ServerClientReset(t *testing.T) { + readErr := make(chan error, 1) + canceled := make(chan struct{}) + p, _ := newHTTP2ClientPeer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, err := r.Body.Read(make([]byte, 1)) + readErr <- err + <-r.Context().Done() + close(canceled) + })) + p.writeConnect(1) + f := p.readFrame() + require.Equal(t, "200", f.(*http2.MetaHeadersFrame).PseudoValue("status")) + require.NoError(t, p.fr.WriteRSTStream(1, http2.ErrCodeCancel)) + select { + case err := <-readErr: + require.Equal(t, http2.StreamError{StreamID: 1, Code: http2.ErrCodeCancel}, err) + case <-time.After(5 * time.Second): + t.Fatal("the body read did not fail after RST_STREAM") + } + select { + case <-canceled: + case <-time.After(5 * time.Second): + t.Fatal("the request context was not canceled") + } +} + +func TestHTTP2ServerAnswersPings(t *testing.T) { + p, _ := newHTTP2ClientPeer(t, http.NotFoundHandler()) + data := [8]byte{8, 7, 6, 5, 4, 3, 2, 1} + require.NoError(t, p.fr.WritePing(false, data)) + f := p.readFrame() + require.IsType(t, &http2.PingFrame{}, f) + require.True(t, f.(*http2.PingFrame).IsAck()) + require.Equal(t, data, f.(*http2.PingFrame).Data) +} + +func TestHTTP2ServerRejectsOverflow(t *testing.T) { + p, _ := newHTTP2ClientPeer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + <-r.Context().Done() + })) + p.writeConnect(1) + p.readFrame() + chunk := make([]byte, http2DefaultFrameSize) + go func() { + for range http2StreamWindow/len(chunk) + 1 { + if p.fr.WriteData(1, false, chunk) != nil { + return + } + } + }() + p.conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + _, err := io.Copy(io.Discard, p.conn) + require.NoError(t, err) +} + +func TestHTTP2ServerBadPreface(t *testing.T) { + conn := serveHTTP2Pipe(t, http.NotFoundHandler()) + _, err := io.WriteString(conn, "GET / HTTP/1.1\r\nHost: example.com\r\n\r\n") + require.NoError(t, err) + conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + n, err := io.Copy(io.Discard, conn) + require.NoError(t, err) + require.Zero(t, n) +} diff --git a/transport/internet/masque/http2_test.go b/transport/internet/masque/http2_test.go new file mode 100644 index 000000000..b2d73b71d --- /dev/null +++ b/transport/internet/masque/http2_test.go @@ -0,0 +1,405 @@ +package masque + +import ( + "bytes" + "context" + "io" + "net" + "net/http" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" + "golang.org/x/net/http2" + "golang.org/x/net/http2/hpack" +) + +type http2Peer struct { + t *testing.T + conn net.Conn + fr *http2.Framer + hbuf bytes.Buffer + henc *hpack.Encoder +} + +func tcpPipe(t *testing.T) (net.Conn, net.Conn) { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + accepted := make(chan net.Conn, 1) + go func() { + conn, _ := ln.Accept() + accepted <- conn + }() + client, err := net.Dial("tcp", ln.Addr().String()) + require.NoError(t, err) + server := <-accepted + require.NotNil(t, server) + return client, server +} + +func newHTTP2Peer(t *testing.T, settings ...http2.Setting) (*http2ClientConn, *http2Peer) { + t.Helper() + client, server := tcpPipe(t) + p := &http2Peer{t: t, conn: server, fr: http2.NewFramer(server, server)} + p.henc = hpack.NewEncoder(&p.hbuf) + p.fr.ReadMetaHeaders = hpack.NewDecoder(4096, nil) + t.Cleanup(func() { server.Close() }) + + ccErr := make(chan error, 1) + var cc *http2ClientConn + go func() { + var err error + cc, err = newHTTP2ClientConn(client) + ccErr <- err + }() + preface := make([]byte, len(http2.ClientPreface)) + _, err := io.ReadFull(server, preface) + require.NoError(t, err) + require.Equal(t, http2.ClientPreface, string(preface)) + + f := p.readFrame() + require.IsType(t, &http2.SettingsFrame{}, f) + var got []http2.Setting + f.(*http2.SettingsFrame).ForeachSetting(func(s http2.Setting) error { + got = append(got, s) + return nil + }) + require.Equal(t, []http2.Setting{ + {ID: http2.SettingHeaderTableSize, Val: http2HeaderTableSize}, + {ID: http2.SettingEnablePush, Val: 0}, + {ID: http2.SettingInitialWindowSize, Val: http2StreamWindow}, + {ID: http2.SettingMaxHeaderListSize, Val: http2MaxHeaderListSize}, + }, got) + f = p.readFrame() + require.IsType(t, &http2.WindowUpdateFrame{}, f) + require.Equal(t, uint32(0), f.Header().StreamID) + require.Equal(t, uint32(http2ConnectionWindow-http2DefaultWindow), f.(*http2.WindowUpdateFrame).Increment) + require.NoError(t, <-ccErr) + t.Cleanup(func() { cc.Close() }) + + require.NoError(t, p.fr.WriteSettings(settings...)) + f = p.readFrame() + require.IsType(t, &http2.SettingsFrame{}, f) + require.True(t, f.(*http2.SettingsFrame).IsAck()) + return cc, p +} + +func (p *http2Peer) readFrame() http2.Frame { + p.t.Helper() + p.conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + f, err := p.fr.ReadFrame() + require.NoError(p.t, err) + return f +} + +func (p *http2Peer) writeHeaders(endStream bool, fields ...string) { + p.t.Helper() + p.hbuf.Reset() + for i := 0; i < len(fields); i += 2 { + require.NoError(p.t, p.henc.WriteField(hpack.HeaderField{Name: fields[i], Value: fields[i+1]})) + } + require.NoError(p.t, p.fr.WriteHeaders(http2.HeadersFrameParam{ + StreamID: http2StreamID, + BlockFragment: p.hbuf.Bytes(), + EndHeaders: true, + EndStream: endStream, + })) +} + +func connectRequest(t *testing.T, ctx context.Context, body io.ReadCloser) *http.Request { + req, err := http.NewRequestWithContext(ctx, http.MethodConnect, "https://proxy.example/.well-known/masque/ip/*/*/", body) + require.NoError(t, err) + req.Header[":protocol"] = []string{"connect-ip"} + req.Header.Set("Capsule-Protocol", "?1") + req.Header.Set("Authorization", "Basic dTpw") + req.Header["User-Agent"] = nil + return req +} + +func TestHTTP2ClientRequest(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + pr, pw := io.Pipe() + + type result struct { + rsp *http.Response + err error + } + results := make(chan result, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), pr)) + results <- result{rsp, err} + }() + + f := p.readFrame() + require.IsType(t, &http2.MetaHeadersFrame{}, f) + headers := f.(*http2.MetaHeadersFrame) + require.False(t, headers.StreamEnded()) + var fields []string + for _, hf := range headers.Fields { + fields = append(fields, hf.Name+": "+hf.Value) + } + require.Equal(t, []string{ + ":method: CONNECT", + ":authority: proxy.example", + ":scheme: https", + ":path: /.well-known/masque/ip/*/*/", + ":protocol: connect-ip", + "authorization: Basic dTpw", + "capsule-protocol: ?1", + }, fields) + + p.writeHeaders(false, ":status", "200", "capsule-protocol", "?1") + r := <-results + require.NoError(t, r.err) + require.Equal(t, http.StatusOK, r.rsp.StatusCode) + require.Equal(t, "?1", r.rsp.Header.Get("Capsule-Protocol")) + + go pw.Write([]byte("ping")) + f = p.readFrame() + require.IsType(t, &http2.DataFrame{}, f) + require.Equal(t, "ping", string(f.(*http2.DataFrame).Data())) + + require.NoError(t, p.fr.WriteData(http2StreamID, false, []byte("pong"))) + b := make([]byte, 16) + n, err := r.rsp.Body.Read(b) + require.NoError(t, err) + require.Equal(t, "pong", string(b[:n])) + + require.NoError(t, pw.Close()) + f = p.readFrame() + require.IsType(t, &http2.DataFrame{}, f) + require.True(t, f.(*http2.DataFrame).StreamEnded()) + + require.NoError(t, p.fr.WriteData(http2StreamID, true, nil)) + _, err = r.rsp.Body.Read(b) + require.ErrorIs(t, err, io.EOF) +} + +func TestHTTP2ClientDefaultUserAgent(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + req := connectRequest(t, context.Background(), nil) + delete(req.Header, "User-Agent") + go cc.RoundTrip(req) + f := p.readFrame() + require.IsType(t, &http2.MetaHeadersFrame{}, f) + var userAgents []string + for _, hf := range f.(*http2.MetaHeadersFrame).Fields { + if hf.Name == "user-agent" { + userAgents = append(userAgents, hf.Value) + } + } + require.Equal(t, []string{http2DefaultUserAgent}, userAgents) +} + +func TestHTTP2ClientNeedsExtendedConnect(t *testing.T) { + cc, _ := newHTTP2Peer(t) + _, err := cc.RoundTrip(connectRequest(t, context.Background(), io.NopCloser(strings.NewReader("")))) + require.ErrorIs(t, err, errHTTP2NoExtendedConnect) +} + +func TestHTTP2ClientSingleStream(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + go cc.RoundTrip(connectRequest(t, context.Background(), nil)) + p.readFrame() + _, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + require.ErrorIs(t, err, errHTTP2StreamUsed) +} + +func TestHTTP2ClientFlowControl(t *testing.T) { + cc, p := newHTTP2Peer(t, + http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}, + http2.Setting{ID: http2.SettingInitialWindowSize, Val: 10}, + ) + pr, pw := io.Pipe() + go cc.RoundTrip(connectRequest(t, context.Background(), pr)) + require.IsType(t, &http2.MetaHeadersFrame{}, p.readFrame()) + + go pw.Write([]byte("0123456789abcdef")) + f := p.readFrame() + require.Equal(t, "0123456789", string(f.(*http2.DataFrame).Data())) + + require.NoError(t, p.fr.WriteWindowUpdate(http2StreamID, 4)) + f = p.readFrame() + require.Equal(t, "abcd", string(f.(*http2.DataFrame).Data())) + + require.NoError(t, p.fr.WriteSettings(http2.Setting{ID: http2.SettingInitialWindowSize, Val: 12})) + var acked bool + var data string + for range 2 { + switch f := p.readFrame().(type) { + case *http2.SettingsFrame: + acked = f.IsAck() + case *http2.DataFrame: + data = string(f.Data()) + } + } + require.True(t, acked) + require.Equal(t, "ef", data) +} + +func TestHTTP2ClientReceiveWindow(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + rsps := make(chan *http.Response, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + if err == nil { + rsps <- rsp + } + }() + require.IsType(t, &http2.MetaHeadersFrame{}, p.readFrame()) + require.True(t, p.readFrame().(*http2.DataFrame).StreamEnded()) + p.writeHeaders(false, ":status", "200") + rsp := <-rsps + + chunk := bytes.Repeat([]byte("x"), http2DefaultFrameSize) + sent := 0 + go func() { + for sent+len(chunk) <= http2WindowUpdateSize { + if p.fr.WriteData(http2StreamID, false, chunk) != nil { + return + } + sent += len(chunk) + } + }() + _, err := io.CopyN(io.Discard, rsp.Body, http2WindowUpdateSize) + require.NoError(t, err) + for _, id := range []uint32{0, http2StreamID} { + f := p.readFrame() + require.IsType(t, &http2.WindowUpdateFrame{}, f) + require.Equal(t, id, f.Header().StreamID) + require.Equal(t, uint32(http2WindowUpdateSize), f.(*http2.WindowUpdateFrame).Increment) + } +} + +func TestHTTP2ClientRejectsOverflow(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + rsps := make(chan *http.Response, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + if err == nil { + rsps <- rsp + } + }() + p.readFrame() + p.readFrame() + p.writeHeaders(false, ":status", "200") + rsp := <-rsps + + chunk := make([]byte, http2DefaultFrameSize) + go func() { + for range http2StreamWindow/len(chunk) + 1 { + if p.fr.WriteData(http2StreamID, false, chunk) != nil { + return + } + } + }() + select { + case <-cc.done: + case <-time.After(5 * time.Second): + t.Fatal("the connection outlived a flow control violation") + } + require.ErrorIs(t, cc.connErr(), http2.ConnectionError(http2.ErrCodeFlowControl)) + _, err := io.Copy(io.Discard, rsp.Body) + require.ErrorIs(t, err, http2.ConnectionError(http2.ErrCodeFlowControl)) +} + +func TestHTTP2ClientRejectsOversizedFrames(t *testing.T) { + cc, p := newHTTP2Peer(t) + require.NoError(t, p.fr.WritePing(false, [8]byte{})) + require.True(t, p.readFrame().(*http2.PingFrame).IsAck()) + + p.fr.AllowIllegalWrites = true + require.NoError(t, p.fr.WriteData(http2StreamID, false, make([]byte, http2DefaultFrameSize+1))) + select { + case <-cc.done: + case <-time.After(5 * time.Second): + t.Fatal("the connection accepted a frame larger than it allows") + } + require.ErrorIs(t, cc.connErr(), http2.ErrFrameTooLarge) +} + +func TestHTTP2ClientStatus(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + rsps := make(chan *http.Response, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + if err == nil { + rsps <- rsp + } + }() + p.readFrame() + p.readFrame() + p.writeHeaders(false, ":status", "100") + p.writeHeaders(true, ":status", "407", "proxy-authenticate", "Basic") + rsp := <-rsps + require.Equal(t, http.StatusProxyAuthRequired, rsp.StatusCode) + require.Equal(t, "Basic", rsp.Header.Get("Proxy-Authenticate")) + _, err := rsp.Body.Read(make([]byte, 1)) + require.ErrorIs(t, err, io.EOF) +} + +func TestHTTP2ClientReset(t *testing.T) { + t.Run("by the server", func(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + errs := make(chan error, 1) + go func() { + _, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + errs <- err + }() + p.readFrame() + p.readFrame() + require.NoError(t, p.fr.WriteRSTStream(http2StreamID, http2.ErrCodeRefusedStream)) + require.Equal(t, http2.StreamError{StreamID: http2StreamID, Code: http2.ErrCodeRefusedStream}, <-errs) + }) + + t.Run("by the context", func(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + ctx, cancel := context.WithCancel(context.Background()) + pr, pw := io.Pipe() + defer pw.Close() + rsps := make(chan *http.Response, 1) + go func() { + rsp, err := cc.RoundTrip(connectRequest(t, ctx, pr)) + if err == nil { + rsps <- rsp + } + }() + p.readFrame() + p.writeHeaders(false, ":status", "200") + rsp := <-rsps + cancel() + f := p.readFrame() + require.IsType(t, &http2.RSTStreamFrame{}, f) + require.Equal(t, http2.ErrCodeCancel, f.(*http2.RSTStreamFrame).ErrCode) + _, err := rsp.Body.Read(make([]byte, 1)) + require.ErrorIs(t, err, context.Canceled) + _, err = pw.Write([]byte("x")) + require.ErrorIs(t, err, io.ErrClosedPipe) + }) + + t.Run("by GOAWAY", func(t *testing.T) { + cc, p := newHTTP2Peer(t, http2.Setting{ID: http2.SettingEnableConnectProtocol, Val: 1}) + errs := make(chan error, 1) + go func() { + _, err := cc.RoundTrip(connectRequest(t, context.Background(), nil)) + errs <- err + }() + p.readFrame() + p.readFrame() + require.NoError(t, p.fr.WriteGoAway(0, http2.ErrCodeNo, nil)) + require.ErrorContains(t, <-errs, "GOAWAY") + }) +} + +func TestHTTP2ClientAnswersPings(t *testing.T) { + _, p := newHTTP2Peer(t) + data := [8]byte{1, 2, 3, 4, 5, 6, 7, 8} + require.NoError(t, p.fr.WritePing(false, data)) + f := p.readFrame() + require.IsType(t, &http2.PingFrame{}, f) + require.True(t, f.(*http2.PingFrame).IsAck()) + require.Equal(t, data, f.(*http2.PingFrame).Data) +} diff --git a/transport/internet/masque/hub.go b/transport/internet/masque/hub.go new file mode 100644 index 000000000..dec37db38 --- /dev/null +++ b/transport/internet/masque/hub.go @@ -0,0 +1,422 @@ +package masque + +import ( + "context" + "crypto/rand" + gotls "crypto/tls" + go_errors "errors" + "io" + "maps" + "net/http" + "net/url" + "runtime" + "slices" + "strings" + "sync" + "time" + + "github.com/apernet/quic-go" + "github.com/apernet/quic-go/http3" + "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/errors" + "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/transport/internet" + "github.com/xtls/xray-core/transport/internet/hysteria/congestion" + "github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr" + "github.com/xtls/xray-core/transport/internet/masque/connectip" + "github.com/xtls/xray-core/transport/internet/tls" + "golang.org/x/net/http2" +) + +type Listener struct { + path pathMatcher + addConn internet.ConnHandler + ctx context.Context + cancel context.CancelFunc + + quicServer *http3.Server + quicListener *quic.Listener + transport *quic.Transport + pktConn net.PacketConn + tcpListener net.Listener + + mu sync.Mutex + conns map[net.Conn]struct{} +} + +func serverVersions(config *tls.Config) (h2, h3 bool) { + h2 = slices.Contains(config.NextProtocol, http2.NextProtoTLS) + h3 = slices.Contains(config.NextProtocol, http3.NextProtoH3) || !h2 + return h2, h3 +} + +func Listen(ctx context.Context, address net.Address, port net.Port, streamSettings *internet.MemoryStreamConfig, handler internet.ConnHandler) (internet.Listener, error) { + if address.Family().IsDomain() { + return nil, errors.New("address is domain") + } + tlsConfig := tls.ConfigFromStreamSettings(streamSettings) + if tlsConfig == nil { + return nil, errors.New("tls config is nil") + } + config := streamSettings.ProtocolSettings.(*Config) + path, err := newPathMatcher(config.Path) + if err != nil { + return nil, err + } + + l := &Listener{ + path: path, + addConn: handler, + conns: make(map[net.Conn]struct{}), + } + l.ctx, l.cancel = context.WithCancel(context.Background()) + h2, h3 := serverVersions(tlsConfig) + if h3 { + if err := l.listenHTTP3(address, port, streamSettings, tlsConfig); err != nil { + l.Close() + return nil, err + } + errors.LogInfo(ctx, "listening UDP for MASQUE over HTTP/3 on ", address, ":", port) + } + if h2 { + if err := l.listenHTTP2(ctx, address, port, streamSettings, tlsConfig); err != nil { + l.Close() + return nil, err + } + errors.LogInfo(ctx, "listening TCP for MASQUE over HTTP/2 on ", address, ":", port) + } + return l, nil +} + +func (l *Listener) listenHTTP3(address net.Address, port net.Port, streamSettings *internet.MemoryStreamConfig, tlsConfig *tls.Config) error { + quicParams := streamSettings.QuicParams + if quicParams == nil { + quicParams = &internet.QuicParams{ + BbrProfile: string(bbr.ProfileStandard), + } + } + switch quicParams.Congestion { + case "", "reno", "bbr", "brutal", "force-brutal": + default: + return errors.New("unknown congestion control: ", quicParams.Congestion) + } + quicConfig := &quic.Config{ + InitialStreamReceiveWindow: quicParams.InitStreamReceiveWindow, + MaxStreamReceiveWindow: quicParams.MaxStreamReceiveWindow, + InitialConnectionReceiveWindow: quicParams.InitConnReceiveWindow, + MaxConnectionReceiveWindow: quicParams.MaxConnReceiveWindow, + MaxIdleTimeout: time.Duration(quicParams.MaxIdleTimeout) * time.Second, + KeepAlivePeriod: time.Duration(quicParams.KeepAlivePeriod) * time.Second, + MaxIncomingStreams: quicParams.MaxIncomingStreams, + InitialPacketSize: initialPacketSize, + DisablePathMTUDiscovery: quicParams.DisablePathMtuDiscovery || (runtime.GOOS != "linux" && runtime.GOOS != "windows" && runtime.GOOS != "darwin"), + EnableDatagrams: true, + DisablePathManager: true, + } + if quicParams.MaxIdleTimeout == 0 { + quicConfig.MaxIdleTimeout = 30 * time.Second + } + + udpAddr := &net.UDPAddr{IP: address.IP(), Port: int(port)} + var err error + if streamSettings.FinalMask != nil { + l.pktConn, err = streamSettings.FinalMask.ListenPacket(context.Background(), udpAddr) + } else { + l.pktConn, err = internet.ListenSystemPacket(context.Background(), udpAddr, streamSettings.SocketSettings) + } + if err != nil { + return errors.New("failed to listen UDP on ", address, ":", port).Base(err) + } + var resetKey *quic.StatelessResetKey + if !quicParams.DisableStatelessReset { + resetKey = &quic.StatelessResetKey{} + common.Must2(rand.Read(resetKey[:])) + } + l.transport = &quic.Transport{Conn: l.pktConn, DisableGSO: quicParams.DisableGSO, StatelessResetKey: resetKey} + + gotlsConfig := tlsConfig.GetTLSConfig() + gotlsConfig.NextProtos = []string{http3.NextProtoH3} + l.quicListener, err = l.transport.Listen(gotlsConfig, quicConfig) + if err != nil { + return err + } + l.quicServer = &http3.Server{ + Handler: l, + EnableDatagrams: true, + ConnContext: func(ctx context.Context, conn *quic.Conn) context.Context { + switch quicParams.Congestion { + case "reno": + case "", "bbr", "brutal": + congestion.UseBBR(conn, bbr.Profile(quicParams.BbrProfile)) + case "force-brutal": + congestion.UseBrutal(conn, quicParams.BrutalUp, quicParams.BrutalDisableLossCompensation) + } + return context.WithValue(ctx, connAddrsKey{}, connAddrs{local: conn.LocalAddr(), remote: conn.RemoteAddr()}) + }, + } + go func() { + if err := l.quicServer.ServeListener(l.quicListener); err != nil && !go_errors.Is(err, quic.ErrServerClosed) && !go_errors.Is(err, http.ErrServerClosed) { + errors.LogErrorInner(context.Background(), err, "failed to serve MASQUE over HTTP/3") + } + }() + return nil +} + +func (l *Listener) listenHTTP2(ctx context.Context, address net.Address, port net.Port, streamSettings *internet.MemoryStreamConfig, tlsConfig *tls.Config) error { + tcpAddr := &net.TCPAddr{IP: address.IP(), Port: int(port)} + var err error + if streamSettings.FinalMask != nil { + l.tcpListener, err = streamSettings.FinalMask.Listen(ctx, tcpAddr) + } else { + l.tcpListener, err = internet.ListenSystem(ctx, tcpAddr, streamSettings.SocketSettings) + } + if err != nil { + return errors.New("failed to listen TCP on ", address, ":", port).Base(err) + } + gotlsConfig := tlsConfig.GetTLSConfig() + gotlsConfig.NextProtos = []string{http2.NextProtoTLS} + go l.acceptHTTP2(gotlsConfig) + return nil +} + +func (l *Listener) acceptHTTP2(config *gotls.Config) { + for { + conn, err := l.tcpListener.Accept() + if err != nil { + if l.ctx.Err() != nil || strings.Contains(err.Error(), "closed") { + return + } + errors.LogWarningInner(context.Background(), err, "failed to accept MASQUE connections") + if strings.Contains(err.Error(), "too many") { + time.Sleep(500 * time.Millisecond) + } + continue + } + go l.serveHTTP2Conn(conn, config) + } +} + +func (l *Listener) serveHTTP2Conn(conn net.Conn, config *gotls.Config) { + tlsConn := tls.Server(conn, config).(*tls.Conn) + if !l.track(tlsConn, true) { + tlsConn.Close() + return + } + defer l.track(tlsConn, false) + ctx, cancel := context.WithTimeout(l.ctx, http2HandshakeTimeout) + err := tlsConn.HandshakeContext(ctx) + cancel() + if err != nil { + errors.LogDebugInner(context.Background(), err, "MASQUE: TLS handshake failed") + tlsConn.Close() + return + } + if protocol := tlsConn.NegotiatedProtocol(); protocol != http2.NextProtoTLS { + errors.LogDebug(context.Background(), "MASQUE: the client negotiated ", protocol, " instead of h2") + tlsConn.Close() + return + } + serveHTTP2(l.ctx, tlsConn, l) +} + +func (l *Listener) track(conn net.Conn, add bool) bool { + l.mu.Lock() + defer l.mu.Unlock() + if !add { + delete(l.conns, conn) + return true + } + if l.ctx.Err() != nil { + return false + } + l.conns[conn] = struct{}{} + return true +} + +func (l *Listener) Addr() net.Addr { + if l.tcpListener != nil { + return l.tcpListener.Addr() + } + return l.quicListener.Addr() +} + +func (l *Listener) Close() error { + l.cancel() + var errs []error + if l.quicServer != nil { + errs = append(errs, l.quicServer.Close()) + } + if l.quicListener != nil { + errs = append(errs, l.quicListener.Close()) + } + if l.transport != nil { + errs = append(errs, l.transport.Close()) + } + if l.pktConn != nil { + errs = append(errs, l.pktConn.Close()) + } + if l.tcpListener != nil { + errs = append(errs, l.tcpListener.Close()) + } + l.mu.Lock() + for conn := range l.conns { + conn.Close() + } + l.mu.Unlock() + return errors.Combine(errs...) +} + +func (l *Listener) ServeHTTP(w http.ResponseWriter, r *http.Request) { + if !l.path.match(r.URL) { + w.WriteHeader(http.StatusNotFound) + return + } + request, err := connectip.ParseProxyRequest(r) + if err != nil { + status := http.StatusBadRequest + if perr, ok := go_errors.AsType[*connectip.ProxyRequestParseError](err); ok { + status = perr.HTTPStatus + } + w.WriteHeader(status) + return + } + addrs, _ := r.Context().Value(connAddrsKey{}).(connAddrs) + conn := &ServerConn{ + w: w, + request: r, + proxyRequest: request, + local: addrs.local, + remote: addrs.remote, + done: make(chan struct{}), + } + l.addConn(conn) + select { + case <-conn.done: + case <-l.ctx.Done(): + conn.Close() + } +} + +type pathMatcher struct { + path string + query url.Values +} + +func newPathMatcher(path string) (pathMatcher, error) { + u, err := url.ParseRequestURI(path) + if err != nil || !strings.HasPrefix(u.Path, "/") { + return pathMatcher{}, errors.New("invalid path: ", path) + } + return pathMatcher{path: u.Path, query: u.Query()}, nil +} + +func (m pathMatcher) match(u *url.URL) bool { + return u.Path == m.path && maps.EqualFunc(u.Query(), m.query, slices.Equal[[]string]) +} + +type ServerConn struct { + w http.ResponseWriter + request *http.Request + proxyRequest *connectip.ProxyRequest + local net.Addr + remote net.Addr + answer sync.Once + mu sync.Mutex + ipConn *connectip.Conn + done chan struct{} + closeOnce sync.Once +} + +func (c *ServerConn) Request() *http.Request { + return c.request +} + +func (c *ServerConn) Accept() (*connectip.Conn, error) { + var err error = errors.New("the request was already answered") + c.answer.Do(func() { + var ipConn *connectip.Conn + ipConn, err = (&connectip.Proxy{}).Proxy(c.w, c.proxyRequest) + c.mu.Lock() + c.ipConn = ipConn + c.mu.Unlock() + }) + if err != nil { + return nil, err + } + return c.ipConn, nil +} + +func (c *ServerConn) Reject(status int, header http.Header) { + c.answer.Do(func() { + for k, vv := range header { + for _, v := range vv { + c.w.Header().Add(k, v) + } + } + c.w.WriteHeader(status) + }) +} + +func (c *ServerConn) tunnel() *connectip.Conn { + c.mu.Lock() + defer c.mu.Unlock() + return c.ipConn +} + +func (c *ServerConn) Read(b []byte) (int, error) { + ipConn := c.tunnel() + if ipConn == nil { + return 0, io.ErrClosedPipe + } + return ipConn.ReadPacket(b) +} + +func (c *ServerConn) Write(b []byte) (int, error) { + ipConn := c.tunnel() + if ipConn == nil { + return 0, io.ErrClosedPipe + } + icmp, err := ipConn.WritePacket(b) + if err != nil { + return 0, err + } + if len(icmp) > 0 { + return 0, &PacketTooBigError{ICMP: icmp} + } + return len(b), nil +} + +func (c *ServerConn) Close() error { + c.closeOnce.Do(func() { + c.Reject(http.StatusInternalServerError, nil) + if ipConn := c.tunnel(); ipConn != nil { + ipConn.Close() + } + close(c.done) + }) + return nil +} + +func (c *ServerConn) LocalAddr() net.Addr { + return c.local +} + +func (c *ServerConn) RemoteAddr() net.Addr { + return c.remote +} + +func (c *ServerConn) SetDeadline(time.Time) error { + return nil +} + +func (c *ServerConn) SetReadDeadline(time.Time) error { + return nil +} + +func (c *ServerConn) SetWriteDeadline(time.Time) error { + return nil +} + +func init() { + common.Must(internet.RegisterTransportListener(protocolName, Listen)) +} diff --git a/transport/internet/masque/hub_test.go b/transport/internet/masque/hub_test.go new file mode 100644 index 000000000..c3069fff3 --- /dev/null +++ b/transport/internet/masque/hub_test.go @@ -0,0 +1,150 @@ +package masque + +import ( + "context" + "io" + "net/http" + "net/http/httptest" + "net/url" + "testing" + "time" + + "github.com/stretchr/testify/require" + "github.com/xtls/xray-core/transport/internet/stat" + "github.com/xtls/xray-core/transport/internet/tls" +) + +func TestServerVersions(t *testing.T) { + for _, c := range []struct { + alpn []string + h2, h3 bool + }{ + {alpn: nil, h3: true}, + {alpn: []string{"h3"}, h3: true}, + {alpn: []string{"h2"}, h2: true}, + {alpn: []string{"h2", "http/1.1"}, h2: true}, + {alpn: []string{"h3", "h2"}, h2: true, h3: true}, + {alpn: []string{"http/1.1"}, h3: true}, + } { + h2, h3 := serverVersions(&tls.Config{NextProtocol: c.alpn}) + if h2 != c.h2 || h3 != c.h3 { + t.Errorf("serverVersions(%q) = %v, %v, want %v, %v", c.alpn, h2, h3, c.h2, c.h3) + } + } +} + +func TestPathMatcher(t *testing.T) { + for _, c := range []struct { + path string + request string + want bool + }{ + {DefaultPath, "/.well-known/masque/ip/*/*/", true}, + {DefaultPath, "/.well-known/masque/ip/%2A/%2A/", true}, + {DefaultPath, "/.well-known/masque/ip/*/*", false}, + {DefaultPath, "/.well-known/masque/ip/*/*/?x=1", false}, + {DefaultPath, "/.well-known/masque/ip/192.0.2.1/6/", false}, + {"/masque?target=*&ipproto=*", "/masque?ipproto=*&target=*", true}, + {"/masque?target=*&ipproto=*", "/masque?target=*", false}, + } { + m, err := newPathMatcher(c.path) + require.NoError(t, err) + u, err := url.ParseRequestURI(c.request) + require.NoError(t, err) + if got := m.match(u); got != c.want { + t.Errorf("path %q matching %q = %v, want %v", c.path, c.request, got, c.want) + } + } + _, err := newPathMatcher("masque") + require.Error(t, err) +} + +func connectIPRequest(target string) *http.Request { + r := httptest.NewRequest(http.MethodGet, "https://proxy.example"+target, nil) + r.Method = http.MethodConnect + r.Proto, r.ProtoMajor, r.ProtoMinor = "HTTP/2.0", 2, 0 + r.Header.Set(":protocol", "connect-ip") + r.Header.Set("Capsule-Protocol", "?1") + r.Body = io.NopCloser(&blockingReader{}) + return r +} + +type blockingReader struct{} + +func (*blockingReader) Read([]byte) (int, error) { + time.Sleep(time.Hour) + return 0, io.EOF +} + +func serve(t *testing.T, r *http.Request, handle func(*ServerConn)) *httptest.ResponseRecorder { + t.Helper() + path, err := newPathMatcher(DefaultPath) + require.NoError(t, err) + l := &Listener{path: path, addConn: func(conn stat.Connection) { + go func() { + handle(conn.(*ServerConn)) + conn.Close() + }() + }} + l.ctx, l.cancel = context.WithCancel(context.Background()) + defer l.cancel() + w := httptest.NewRecorder() + done := make(chan struct{}) + go func() { + l.ServeHTTP(w, r) + close(done) + }() + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("ServeHTTP did not return") + } + return w +} + +func TestListenerRejectsOtherRequests(t *testing.T) { + unexpected := func(*ServerConn) { t.Error("an invalid request reached the proxy") } + + require.Equal(t, http.StatusNotFound, serve(t, connectIPRequest("/other"), unexpected).Code) + + get := connectIPRequest(DefaultPath) + get.Method = http.MethodGet + require.Equal(t, http.StatusMethodNotAllowed, serve(t, get, unexpected).Code) + + websocket := connectIPRequest(DefaultPath) + websocket.Header.Set(":protocol", "websocket") + require.Equal(t, http.StatusNotImplemented, serve(t, websocket, unexpected).Code) + + noCapsules := connectIPRequest(DefaultPath) + noCapsules.Header.Del("Capsule-Protocol") + require.Equal(t, http.StatusBadRequest, serve(t, noCapsules, unexpected).Code) +} + +func TestServerConnAnswers(t *testing.T) { + t.Run("reject", func(t *testing.T) { + w := serve(t, connectIPRequest(DefaultPath), func(conn *ServerConn) { + require.Equal(t, "connect-ip", conn.Request().Header.Get(":protocol")) + conn.Reject(http.StatusUnauthorized, http.Header{"WWW-Authenticate": {"Basic"}}) + }) + require.Equal(t, http.StatusUnauthorized, w.Code) + require.Equal(t, "Basic", w.Header().Get("WWW-Authenticate")) + }) + + t.Run("no answer", func(t *testing.T) { + w := serve(t, connectIPRequest(DefaultPath), func(*ServerConn) {}) + require.Equal(t, http.StatusInternalServerError, w.Code) + }) + + t.Run("accept", func(t *testing.T) { + w := serve(t, connectIPRequest(DefaultPath), func(conn *ServerConn) { + ipConn, err := conn.Accept() + require.NoError(t, err) + require.NotNil(t, ipConn) + _, err = conn.Accept() + require.Error(t, err) + conn.Reject(http.StatusForbidden, nil) + }) + require.Equal(t, http.StatusOK, w.Code) + require.Equal(t, "?1", w.Header().Get("Capsule-Protocol")) + }) +} diff --git a/transport/internet/memory_settings.go b/transport/internet/memory_settings.go index 770cf82cc..6277151c6 100644 --- a/transport/internet/memory_settings.go +++ b/transport/internet/memory_settings.go @@ -54,11 +54,10 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) { mss.SecurityType = s.SecurityType mss.SecuritySettings = ess } + if s != nil && (len(s.Tcpmasks) != 0 || len(s.Udpmasks) != 0) { + var tcpMasks []finalmask.TCPMask + var udpMasks []finalmask.UDPMask - var tcpMasks []finalmask.TCPMask - var udpMasks []finalmask.UDPMask - - if s != nil { for i := range s.Tcpmasks { instance := common.Must2(s.Tcpmasks[i].GetInstance()) tcpMasks = append(tcpMasks, instance.(finalmask.TCPMask)) @@ -67,37 +66,37 @@ func ToMemoryStreamConfig(s *StreamConfig) (*MemoryStreamConfig, error) { instance := common.Must2(s.Udpmasks[i].GetInstance()) udpMasks = append(udpMasks, instance.(finalmask.UDPMask)) } - } - dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { - return DialSystem(ctx, dest, mss.SocketSettings) - } - listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) { - return ListenSystem(ctx, addr, mss.SocketSettings) - } - dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) { - conn, err := DialSystem(ctx, dest, mss.SocketSettings) - if err != nil { - return nil, nil, err + dialTCP := func(ctx context.Context, dest net.Destination) (net.Conn, error) { + return DialSystem(ctx, dest, mss.SocketSettings) } - var newConn net.PacketConn - var udpAddr net.Addr - switch c := conn.(type) { - case *PacketConnWrapper: - newConn = c.PacketConn - udpAddr = conn.RemoteAddr() - case *cnc.Connection: - newConn = &FakePacketConn{Conn: c} - udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0} - default: - panic(reflect.TypeOf(c)) + listen := func(ctx context.Context, addr net.Addr) (net.Listener, error) { + return ListenSystem(ctx, addr, mss.SocketSettings) } - return newConn, udpAddr, nil + dialUDP := func(ctx context.Context, dest net.Destination) (net.PacketConn, net.Addr, error) { + conn, err := DialSystem(ctx, dest, mss.SocketSettings) + if err != nil { + return nil, nil, err + } + var newConn net.PacketConn + var udpAddr net.Addr + switch c := conn.(type) { + case *net.PacketConnWrapper: + newConn = c.PacketConn + udpAddr = conn.RemoteAddr() + case *cnc.Connection: + newConn = &FakePacketConn{Conn: c} + udpAddr = &net.UDPAddr{IP: []byte{0, 0, 0, 0}, Port: 0} + default: + panic(reflect.TypeOf(c)) + } + return newConn, udpAddr, nil + } + listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) { + return ListenSystemPacket(ctx, addr, mss.SocketSettings) + } + mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket) } - listenPacket := func(ctx context.Context, addr net.Addr) (net.PacketConn, error) { - return ListenSystemPacket(ctx, addr, mss.SocketSettings) - } - mss.FinalMask = finalmask.NewFinalMask(tcpMasks, udpMasks, dialTCP, listen, dialUDP, listenPacket) if s != nil && s.QuicParams != nil { mss.QuicParams = s.QuicParams diff --git a/transport/internet/reality/reality.go b/transport/internet/reality/reality.go index 50b2e02f7..9ccae6686 100644 --- a/transport/internet/reality/reality.go +++ b/transport/internet/reality/reality.go @@ -18,6 +18,7 @@ import ( "regexp" "strings" "sync" + "sync/atomic" "time" "unsafe" @@ -36,6 +37,18 @@ import ( type Conn struct { *reality.Conn + suppressCloseNotify atomic.Bool +} + +func (c *Conn) SuppressCloseNotify() { + c.suppressCloseNotify.Store(true) +} + +func (c *Conn) Close() error { + if c.suppressCloseNotify.Load() { + return c.Conn.NetConn().Close() + } + return c.Conn.Close() } func (c *Conn) HandshakeAddress() net.Address { @@ -56,10 +69,22 @@ func Server(c net.Conn, config *reality.Config) (net.Conn, error) { type UConn struct { *utls.UConn - Config *Config - ServerName string - AuthKey []byte - Verified bool + Config *Config + ServerName string + AuthKey []byte + Verified bool + suppressCloseNotify atomic.Bool +} + +func (c *UConn) SuppressCloseNotify() { + c.suppressCloseNotify.Store(true) +} + +func (c *UConn) Close() error { + if c.suppressCloseNotify.Load() { + return c.NetConn().Close() + } + return c.UConn.Close() } func (c *UConn) HandshakeAddress() net.Address { @@ -132,7 +157,7 @@ func UClient(c net.Conn, config *Config, ctx context.Context, dest net.Destinati uConn.ServerName = utlsConfig.ServerName fingerprint := tls.GetFingerprint(config.Fingerprint) if fingerprint == nil { - return nil, errors.New("REALITY: failed to get fingerprint").AtError() + return nil, errors.New("REALITY: failed to get fingerprint") } uConn.UConn = utls.UClient(c, utlsConfig, *fingerprint) { @@ -271,7 +296,7 @@ func UClient(c net.Conn, config *Config, ctx context.Context, dest net.Destinati // Do not close the connection }() time.Sleep(time.Duration(crypto.RandBetween(config.SpiderY[8], config.SpiderY[9])) * time.Millisecond) // return - return nil, errors.New("REALITY: processed invalid connection").AtWarning() + return nil, errors.New("REALITY: processed invalid connection") } return uConn, nil } diff --git a/transport/internet/sockopt_darwin.go b/transport/internet/sockopt_darwin.go index 73a85d8ff..2ac78ab4a 100644 --- a/transport/internet/sockopt_darwin.go +++ b/transport/internet/sockopt_darwin.go @@ -288,14 +288,14 @@ func applyInboundSocketOptions(network string, fd uintptr, config *SocketConfig) func setReuseAddr(fd uintptr) error { if err := unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEADDR, 1); err != nil { - return errors.New("failed to set SO_REUSEADDR").Base(err).AtWarning() + return errors.New("failed to set SO_REUSEADDR").Base(err) } return nil } func setReusePort(fd uintptr) error { if err := unix.SetsockoptInt(int(fd), unix.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil { - return errors.New("failed to set SO_REUSEPORT").Base(err).AtWarning() + return errors.New("failed to set SO_REUSEPORT").Base(err) } return nil } diff --git a/transport/internet/sockopt_freebsd.go b/transport/internet/sockopt_freebsd.go index 635c25da3..f373408d1 100644 --- a/transport/internet/sockopt_freebsd.go +++ b/transport/internet/sockopt_freebsd.go @@ -224,7 +224,7 @@ func applyInboundSocketOptions(network string, fd uintptr, config *SocketConfig) func setReuseAddr(fd uintptr) error { if err := syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, syscall.SO_REUSEADDR, 1); err != nil { - return errors.New("failed to set SO_REUSEADDR").Base(err).AtWarning() + return errors.New("failed to set SO_REUSEADDR").Base(err) } return nil } @@ -232,7 +232,7 @@ func setReuseAddr(fd uintptr) error { func setReusePort(fd uintptr) error { if err := syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, soReUsePortLB, 1); err != nil { if err := syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, soReUsePort, 1); err != nil { - return errors.New("failed to set SO_REUSEPORT").Base(err).AtWarning() + return errors.New("failed to set SO_REUSEPORT").Base(err) } } return nil diff --git a/transport/internet/sockopt_linux.go b/transport/internet/sockopt_linux.go index 36c4decfd..7c614d9fb 100644 --- a/transport/internet/sockopt_linux.go +++ b/transport/internet/sockopt_linux.go @@ -234,14 +234,14 @@ func applyInboundSocketOptions(network string, fd uintptr, config *SocketConfig) func setReuseAddr(fd uintptr) error { if err := syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, syscall.SO_REUSEADDR, 1); err != nil { - return errors.New("failed to set SO_REUSEADDR").Base(err).AtWarning() + return errors.New("failed to set SO_REUSEADDR").Base(err) } return nil } func setReusePort(fd uintptr) error { if err := syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, unix.SO_REUSEPORT, 1); err != nil { - return errors.New("failed to set SO_REUSEPORT").Base(err).AtWarning() + return errors.New("failed to set SO_REUSEPORT").Base(err) } return nil } diff --git a/transport/internet/splithttp/dialer.go b/transport/internet/splithttp/dialer.go index e896516ed..b52f5a71d 100644 --- a/transport/internet/splithttp/dialer.go +++ b/transport/internet/splithttp/dialer.go @@ -25,7 +25,6 @@ import ( "github.com/xtls/xray-core/common/signal/done" "github.com/xtls/xray-core/transport/internet" "github.com/xtls/xray-core/transport/internet/browser_dialer" - "github.com/xtls/xray-core/transport/internet/finalmask" "github.com/xtls/xray-core/transport/internet/hysteria/congestion" "github.com/xtls/xray-core/transport/internet/hysteria/congestion/bbr" "github.com/xtls/xray-core/transport/internet/reality" @@ -200,7 +199,7 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea if err != nil { return nil, errors.New("failed to dial to dest").Base(err) } - pktConn = conn.(*finalmask.PacketConnWrapper).PacketConn + pktConn = conn.(*net.PacketConnWrapper).PacketConn udpAddr = conn.RemoteAddr() } else { conn, err := internet.DialSystem(ctx, dest, streamSettings.SocketSettings) @@ -208,7 +207,7 @@ func createHTTPClient(dest net.Destination, streamSettings *internet.MemoryStrea return nil, errors.New("failed to dial to dest").Base(err) } switch c := conn.(type) { - case *internet.PacketConnWrapper: + case *net.PacketConnWrapper: pktConn = c.PacketConn udpAddr = c.RemoteAddr() case *cnc.Connection: diff --git a/transport/internet/system_dialer.go b/transport/internet/system_dialer.go index 2ff7693de..bfe3fb592 100644 --- a/transport/internet/system_dialer.go +++ b/transport/internet/system_dialer.go @@ -86,7 +86,7 @@ func (d *DefaultSystemDialer) Dial(ctx context.Context, src net.Address, dest ne if err != nil { return nil, err } - return &PacketConnWrapper{ + return &net.PacketConnWrapper{ PacketConn: packetConn, Dest: destAddr, }, nil @@ -148,24 +148,6 @@ func (d *DefaultSystemDialer) DestIpAddress() net.IP { return nil } -type PacketConnWrapper struct { - net.PacketConn - Dest net.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() net.Addr { - return c.Dest -} - type SystemDialerAdapter interface { Dial(network string, address string) (net.Conn, error) } diff --git a/transport/internet/tcp/dialer.go b/transport/internet/tcp/dialer.go index 680ee83b2..47779c744 100644 --- a/transport/internet/tcp/dialer.go +++ b/transport/internet/tcp/dialer.go @@ -83,14 +83,14 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me } if err != nil { if isFromMitmVerify { - return nil, errors.New("MITM freedom RAW TLS: failed to verify Domain Fronting certificate from " + mitmServerName).Base(err).AtWarning() + return nil, errors.New("MITM freedom RAW TLS: failed to verify Domain Fronting certificate from " + mitmServerName).Base(err) } return nil, err } negotiatedProtocol := conn.(tls.Interface).NegotiatedProtocol() if isFromMitmAlpn && !mitmAlpn11 && negotiatedProtocol != "h2" { conn.Close() - return nil, errors.New("MITM freedom RAW TLS: unexpected Negotiated Protocol (" + negotiatedProtocol + ") with " + mitmServerName).AtWarning() + return nil, errors.New("MITM freedom RAW TLS: unexpected Negotiated Protocol (" + negotiatedProtocol + ") with " + mitmServerName) } } else if config := reality.ConfigFromStreamSettings(streamSettings); config != nil { if conn, err = reality.UClient(conn, config, ctx, dest); err != nil { @@ -102,11 +102,11 @@ func Dial(ctx context.Context, dest net.Destination, streamSettings *internet.Me if tcpSettings.HeaderSettings != nil { headerConfig, err := tcpSettings.HeaderSettings.GetInstance() if err != nil { - return nil, errors.New("failed to get header settings").Base(err).AtError() + return nil, errors.New("failed to get header settings").Base(err) } auth, err := internet.CreateConnectionAuthenticator(headerConfig) if err != nil { - return nil, errors.New("failed to create header authenticator").Base(err).AtError() + return nil, errors.New("failed to create header authenticator").Base(err) } conn = auth.Client(conn) } diff --git a/transport/internet/tcp/hub.go b/transport/internet/tcp/hub.go index bb8099ed3..92eb6e63f 100644 --- a/transport/internet/tcp/hub.go +++ b/transport/internet/tcp/hub.go @@ -74,11 +74,11 @@ func ListenTCP(ctx context.Context, address net.Address, port net.Port, streamSe if tcpSettings.HeaderSettings != nil { headerConfig, err := tcpSettings.HeaderSettings.GetInstance() if err != nil { - return nil, errors.New("invalid header settings").Base(err).AtError() + return nil, errors.New("invalid header settings").Base(err) } auth, err := internet.CreateConnectionAuthenticator(headerConfig) if err != nil { - return nil, errors.New("invalid header settings.").Base(err).AtError() + return nil, errors.New("invalid header settings.").Base(err) } l.authConfig = auth } diff --git a/transport/internet/tcp_hub.go b/transport/internet/tcp_hub.go index 31183368a..1f613d0e5 100644 --- a/transport/internet/tcp_hub.go +++ b/transport/internet/tcp_hub.go @@ -12,7 +12,7 @@ var transportListenerCache = make(map[string]ListenFunc) func RegisterTransportListener(protocol string, listener ListenFunc) error { if _, found := transportListenerCache[protocol]; found { - return errors.New(protocol, " listener already registered.").AtError() + return errors.New(protocol, " listener already registered.") } transportListenerCache[protocol] = listener return nil @@ -40,7 +40,7 @@ func ListenUnix(ctx context.Context, address net.Address, settings *MemoryStream protocol := settings.ProtocolName listenFunc := transportListenerCache[protocol] if listenFunc == nil { - return nil, errors.New(protocol, " unix listener not registered.").AtError() + return nil, errors.New(protocol, " unix listener not registered.") } listener, err := listenFunc(ctx, address, net.Port(0), settings, handler) if err != nil { @@ -72,7 +72,7 @@ func ListenTCP(ctx context.Context, address net.Address, port net.Port, settings protocol := settings.ProtocolName listenFunc := transportListenerCache[protocol] if listenFunc == nil { - return nil, errors.New(protocol, " listener not registered.").AtError() + return nil, errors.New(protocol, " listener not registered.") } listener, err := listenFunc(ctx, address, port, settings, handler) if err != nil { diff --git a/transport/internet/tls/config.go b/transport/internet/tls/config.go index 9038cd281..f33e829c1 100644 --- a/transport/internet/tls/config.go +++ b/transport/internet/tls/config.go @@ -39,7 +39,7 @@ func (c *Config) loadSelfCertPool() (*x509.CertPool, error) { root := x509.NewCertPool() for _, cert := range c.Certificate { if !root.AppendCertsFromPEM(cert.Certificate) { - return nil, errors.New("failed to append cert").AtWarning() + return nil, errors.New("failed to append cert") } } return root, nil diff --git a/transport/internet/tls/config_other.go b/transport/internet/tls/config_other.go index efd18c933..994723bfe 100644 --- a/transport/internet/tls/config_other.go +++ b/transport/internet/tls/config_other.go @@ -44,11 +44,11 @@ func (c *Config) getCertPool() (*x509.CertPool, error) { pool, err := x509.SystemCertPool() if err != nil { - return nil, errors.New("system root").AtWarning().Base(err) + return nil, errors.New("system root").Base(err) } for _, cert := range c.Certificate { if !pool.AppendCertsFromPEM(cert.Certificate) { - return nil, errors.New("append cert to root").AtWarning().Base(err) + return nil, errors.New("append cert to root").Base(err) } } return pool, nil diff --git a/transport/internet/tls/tls.go b/transport/internet/tls/tls.go index df5d1cbd7..f27fa17b0 100644 --- a/transport/internet/tls/tls.go +++ b/transport/internet/tls/tls.go @@ -6,6 +6,7 @@ import ( "crypto/tls" "math/big" "slices" + "sync/atomic" "time" utls "github.com/refraction-networking/utls" @@ -29,11 +30,19 @@ var ( type Conn struct { *tls.Conn + suppressCloseNotify atomic.Bool } const tlsCloseTimeout = 250 * time.Millisecond +func (c *Conn) SuppressCloseNotify() { + c.suppressCloseNotify.Store(true) +} + func (c *Conn) Close() error { + if c.suppressCloseNotify.Load() { + return c.Conn.NetConn().Close() + } timer := time.AfterFunc(tlsCloseTimeout, func() { c.Conn.NetConn().Close() }) @@ -74,11 +83,19 @@ func Server(c net.Conn, config *tls.Config) net.Conn { type UConn struct { *utls.UConn + suppressCloseNotify atomic.Bool } var _ Interface = (*UConn)(nil) +func (c *UConn) SuppressCloseNotify() { + c.suppressCloseNotify.Store(true) +} + func (c *UConn) Close() error { + if c.suppressCloseNotify.Load() { + return c.Conn.NetConn().Close() + } timer := time.AfterFunc(tlsCloseTimeout, func() { c.Conn.NetConn().Close() })