Compare commits

..
Author SHA1 Message Date
Fangliding 7f23673023 Do not create Finalmask if not needed 2026-09-29 15:25:46 +08:00
Fangliding 3a9412c128 fmt 2026-09-29 15:25:45 +08:00
Fangliding 6140ff6844 Move PacketConnWrapper to common/net 2026-09-29 15:25:40 +08:00
e5e85ca9da XTLS Vision: Suppress outer CloseNotify after switching to direct copy (#6834)
https://github.com/XTLS/Xray-core/pull/6816#issuecomment-5828082031
https://github.com/XTLS/Xray-core/issues/6794#issuecomment-5847200576
https://github.com/XTLS/Xray-core/pull/6834#issuecomment-5860728689

Fixes https://github.com/XTLS/Xray-core/issues/6794#issuecomment-5754894283
Fixes https://github.com/XTLS/Xray-core/issues/6124#issuecomment-4439918073
Fixes https://github.com/XTLS/Xray-core/issues/4878#issuecomment-5754638401

---------

Co-authored-by: Artem Lytkin <146867384+4RH1T3CT0R7@users.noreply.github.com>
2026-09-27 23:39:21 +00:00
7780db9bbe Finalmask: Fix panic when udpHop or xicmp fails; Restore udpHop's default interval (#6808)
https://github.com/XTLS/Xray-core/pull/6808#pullrequestreview-5324529175

---------

Co-authored-by: 风扇滑翔翼 <Fangliding.fshxy@outlook.com>
2026-09-27 22:17:47 +00:00
Artem LytkinandGitHub 2953d44734 Geodata: Reduce MPH matcher runtime memory usage (#6821)
https://github.com/XTLS/Xray-core/pull/6821#issuecomment-5860086534
2026-09-27 21:52:37 +00:00
7b8ade3ec5 TUN inbound: Only close the UDP connection that actually finished (#6814)
https://github.com/XTLS/Xray-core/pull/6814#pullrequestreview-5324536423

---------

Co-authored-by: 风扇滑翔翼 <Fangliding.fshxy@outlook.com>
2026-09-27 21:06:19 +00:00
Artem LytkinandGitHub 5dda894e29 TUN inbound: Reuse existing Wintun adapter by name on Windows again (#6811)
Fixes https://github.com/XTLS/Xray-core/issues/6780
2026-09-27 20:39:17 +00:00
youugiuhiuhandGitHub 47a2c2ffdc TUN inbound: Add autoSystemDNS on Linux (to TUN's gateway) (#6773)
https://github.com/XTLS/Xray-core/issues/6454#issuecomment-4931311976
https://github.com/XTLS/Xray-core/pull/6773#issuecomment-5755516423
https://github.com/XTLS/Xray-core/pull/6807#issuecomment-5807844210
2026-09-27 20:07:53 +00:00
CluvexandGitHub 7a018833ec Proxy: Add MASQUE inbound (IETF CONNECT-IP server, RFC 9484) (#6844)
Completes https://github.com/XTLS/Xray-core/pull/6807 and https://github.com/XTLS/Xray-core/pull/6810
2026-09-27 19:16:54 +00:00
风扇滑翔翼andGitHub 65e853ed84 SS2022 proxy: Refactor to remove sing* dependencies (#6831)
https://github.com/XTLS/Xray-core/pull/4356#issuecomment-2639571931
2026-09-27 18:42:53 +00:00
HuskyDGandGitHub 3519dfecbd FakeDNS: Change default FakeIPv6Pool to 2001:2::/48 (#6815)
https://github.com/XTLS/Xray-core/discussions/5458#discussioncomment-18555021

https://github.com/XTLS/Xray-core/pull/6815#issuecomment-5842712128
2026-09-26 03:52:03 +00:00
CluvexandRPRX df261e4479 MASQUE client: Support HTTP/2 (Extended CONNECT, RFC 8441) (#6810)
https://github.com/XTLS/Xray-core/pull/6807#issuecomment-5808933074

https://github.com/XTLS/Xray-core/pull/6810#issuecomment-5842441136
2026-09-26 03:49:37 +00:00
风扇滑翔翼andRPRX 61cad5ec8b Xray-core: Reduce error log usage for performance (#6796)
https://github.com/XTLS/Xray-core/pull/6796#issuecomment-5807809588
2026-09-25 17:59:49 +00:00
Artem LytkinandGitHub a642a190ed Geodata: Prefilter regexp rules by their required literals (#6818)
https://github.com/XTLS/Xray-core/pull/6818#issuecomment-5833512067
2026-09-25 17:25:52 +00:00
Hossin AsaadiandGitHub 60e2a0c502 WireGuard proxy: Release packet views after use (#6801)
https://github.com/XTLS/Xray-core/pull/6801#issuecomment-5807228464
2026-09-24 05:46:26 +00:00
Esko MobiusandGitHub 7d3e44fee2 Proxy: Add MASQUE outbound & transport (IETF CONNECT-IP, RFC 9484) (#6807)
Closes https://github.com/XTLS/Xray-core/issues/5495#issuecomment-3710683679
2026-09-24 02:13:37 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
9927942aaa Bump google.golang.org/grpc from 1.83.2 to 1.84.0 (#6793)
Bumps [google.golang.org/grpc](https://github.com/grpc/grpc-go) from 1.83.2 to 1.84.0.
- [Release notes](https://github.com/grpc/grpc-go/releases)
- [Commits](https://github.com/grpc/grpc-go/compare/v1.83.2...v1.84.0)

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

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

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

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

---------

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

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

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

---------

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

---------

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

---------

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

---------

Co-authored-by: 风扇滑翔翼 <Fangliding.fshxy@outlook.com>
2026-09-12 19:51:59 +00:00
RPRXandGitHub 52a412d9e2 Xray-core v26.9.9
Sponsor & Donation & NFTs: https://github.com/XTLS/Xray-core/issues/3668
Project X Channel: https://t.me/projectXtls

Announcement of NFTs by Project X: https://github.com/XTLS/Xray-core/discussions/3633
Project X NFT: https://opensea.io/assets/ethereum/0x5ee362866001613093361eb8569d59c4141b76d1/1

VLESS Post-Quantum Encryption: https://github.com/XTLS/Xray-core/pull/5067
VLESS NFT: https://opensea.io/collection/vless

XHTTP: Beyond REALITY: https://github.com/XTLS/Xray-core/discussions/4113
REALITY NFT: https://opensea.io/assets/ethereum/0x5ee362866001613093361eb8569d59c4141b76d1/2
2026-09-08 22:25:43 +00:00
18a1b5042a Finalmask: Make "udpHop" a new UDP mask (#6327)
https://github.com/XTLS/Xray-core/pull/6327#issuecomment-5592421364

---------

Co-authored-by: RPRX <63339210+RPRX@users.noreply.github.com>
2026-09-08 22:15:02 +00:00
patternihaandGitHub c26d2eda24 Direct/Freedom outbound: Skip domain resolution and finalRules when using sockopt.dialerProxy (#6742)
https://github.com/XTLS/Xray-core/pull/6058#issuecomment-5578800461
2026-09-08 18:45:38 +00:00
风扇滑翔翼andGitHub a1bf968be9 VLESS config: Fix validateOutboundTransportSecurity() (#6741)
Fixes https://github.com/XTLS/Xray-core/issues/6737#issuecomment-5587270250
2026-09-08 18:18:12 +00:00
风扇滑翔翼andGitHub c037ccd98d infra/vformat/main.go: Refactor (#6739)
https://github.com/XTLS/Xray-core/pull/6327#issuecomment-5581733848
2026-09-08 17:19:15 +00:00
236 changed files with 22537 additions and 3367 deletions
+1 -3
View File
@@ -67,9 +67,7 @@ jobs:
check-latest: true
cache: false
- name: Check Format
run: |
go install -v mvdan.cc/gofumpt@latest
go run ./infra/vformat/main.go -mode check -pwd ./
run: go run ./infra/vformat/main.go -mode check -pwd ./
test:
needs: check-assets
+1 -1
View File
@@ -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) {
+1 -1
View File
@@ -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
+22
View File
@@ -212,6 +212,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
+59
View File
@@ -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)
}
})
}
}
+2 -2
View File
@@ -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()
+3 -3
View File
@@ -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
+4 -4
View File
@@ -84,7 +84,7 @@ func NewServer(ctx context.Context, dest net.Destination, dispatcher routing.Dis
if dest.Network == net.Network_UDP { // UDP classic DNS mode
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.
@@ -102,7 +102,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)
@@ -113,7 +113,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)
}
}
@@ -122,7 +122,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)
}
}
+2 -2
View File
@@ -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)
+6 -2
View File
@@ -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")
+1 -1
View File
@@ -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})
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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)
+2 -2
View File
@@ -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{
+11 -1
View File
@@ -5,7 +5,8 @@ import (
)
type windowsReader struct {
bufs []syscall.WSABuf
bufs []syscall.WSABuf
ready bool
}
func (r *windowsReader) Init(bs []*Buffer) {
@@ -15,6 +16,7 @@ func (r *windowsReader) Init(bs []*Buffer) {
for _, b := range bs {
r.bufs = append(r.bufs, syscall.WSABuf{Len: uint32(Size), Buf: &b.v[0]})
}
r.ready = false
}
func (r *windowsReader) Clear() {
@@ -25,6 +27,14 @@ func (r *windowsReader) Clear() {
}
func (r *windowsReader) Read(fd uintptr) int32 {
// On the first invocation, we return -1 to indicate "not ready"
// to make rawConn.Read wait for readability using the runtime's own mechanism
// because syscall.WSARecv() is a blocking call when used with nil OVERLAPPED
if !r.ready {
r.ready = true
return -1
}
var nBytes uint32
var flags uint32
err := syscall.WSARecv(syscall.Handle(fd), &r.bufs[0], uint32(len(r.bufs)), &nBytes, &flags, nil, nil)
+3 -3
View File
@@ -10,12 +10,12 @@ import (
// [,)
func RandBetween(from int64, to int64) int64 {
if from == to {
return from
}
if from > to {
from, to = to, from
}
if d := to - from; d == 0 || d == 1 {
return from
}
bigInt, _ := rand.Int(rand.Reader, big.NewInt(to-from))
return from + bigInt.Int64()
}
+13 -65
View File
@@ -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
}
+6 -15
View File
@@ -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)
}
@@ -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) {
+47 -14
View File
@@ -1,6 +1,8 @@
package strmatcher
import (
"errors"
"math"
"math/bits"
"runtime"
"sort"
@@ -38,19 +40,21 @@ type mphRuleInfo struct {
// MphMatcherGroup is an implementation of MatcherGroup.
// It implements Rabin-Karp algorithm and minimal perfect hash table for Full and Domain matcher.
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<seed> & Mask -> stored index for rules
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
ruleInfos *map[string]mphRuleInfo
patterns string // All rule patterns concatenated
patternOffs []uint32 // RuleIdx -> patterns[patternOffs[i]:patternOffs[i+1]], index 0 reserved for failed lookup
values []uint32 // All registered matcher values concatenated
valueOffs []uint32 // RuleIdx -> values[valueOffs[i]:valueOffs[i+1]] (Full Matcher takes precedence)
level0 []uint32 // RollingHash & Mask -> seed for Memhash
level0Mask uint32 // Mask restricting RollingHash to 0 ~ len(level0)
level1 []uint32 // Memhash<seed> & Mask -> stored index for rules
level1Mask uint32 // Mask for restricting Memhash<seed> to 0 ~ len(level1)
rules []string // RuleIdx -> pattern string, only used for building
ruleInfos *map[string]mphRuleInfo
}
func NewMphMatcherGroup() *MphMatcherGroup {
return &MphMatcherGroup{
rules: []string{""},
values: [][]uint32{nil},
level0: nil,
level0Mask: 0,
level1: nil,
@@ -78,7 +82,6 @@ func (g *MphMatcherGroup) addPattern(suffixHash uint32, suffixPattern string, pa
if !found {
info = mphRuleInfo{rollingHash: RollingHash(suffixHash, pattern)}
g.rules = append(g.rules, fullPattern)
g.values = append(g.values, nil)
}
info.matchers[matcherType] = append(info.matchers[matcherType], value)
(*g.ruleInfos)[fullPattern] = info
@@ -94,14 +97,30 @@ func (g *MphMatcherGroup) Build() error {
g.level1 = make([]uint32, nextPow2(ruleCount))
g.level1Mask = uint32(len(g.level1) - 1)
// Flatten patterns and values so the built group has no per-rule objects
valueCount := 0
for _, ruleInfo := range *g.ruleInfos {
valueCount += len(ruleInfo.matchers[Full]) + len(ruleInfo.matchers[Domain])
}
g.patterns = strings.Join(g.rules, "")
if uint64(len(g.patterns)) > math.MaxUint32 || uint64(valueCount) > math.MaxUint32 {
return errors.New("too many rules for MphMatcherGroup")
}
g.patternOffs = make([]uint32, len(g.rules)+1)
g.values = make([]uint32, 0, valueCount)
g.valueOffs = make([]uint32, len(g.rules)+1)
// Create buckets based on all rule's rolling hash
buckets := make([][]uint32, len(g.level0))
for ruleIdx := 1; ruleIdx < len(g.rules); ruleIdx++ { // Traverse rules starting from index 1 (0 reserved for failed lookup)
ruleInfo := (*g.ruleInfos)[g.rules[ruleIdx]]
bucketIdx := ruleInfo.rollingHash & g.level0Mask
buckets[bucketIdx] = append(buckets[bucketIdx], uint32(ruleIdx))
g.values[ruleIdx] = append(ruleInfo.matchers[Full], ruleInfo.matchers[Domain]...) // nolint:gocritic
g.patternOffs[ruleIdx+1] = g.patternOffs[ruleIdx] + uint32(len(g.rules[ruleIdx]))
g.values = append(append(g.values, ruleInfo.matchers[Full]...), ruleInfo.matchers[Domain]...)
g.valueOffs[ruleIdx+1] = uint32(len(g.values))
}
g.rules = nil
g.ruleInfos = nil // Set ruleInfos nil to release memory
runtime.GC() // peak mem
@@ -121,7 +140,7 @@ func (g *MphMatcherGroup) Build() error {
seed := uint32(0)
for len(hashedBucket) != len(bucket) {
for _, ruleIdx := range bucket {
memHash := MemHash(seed, g.rules[ruleIdx]) & g.level1Mask
memHash := MemHash(seed, g.pattern(ruleIdx)) & g.level1Mask
if occupied[memHash] { // Collision occurred with this seed
for _, hash := range hashedBucket { // Revert all values in this hashed bucket
occupied[hash] = false
@@ -141,12 +160,26 @@ func (g *MphMatcherGroup) Build() error {
return nil
}
func (g *MphMatcherGroup) pattern(ruleIdx uint32) string {
return g.patterns[g.patternOffs[ruleIdx]:g.patternOffs[ruleIdx+1]]
}
// valuesOf caps the capacity, so appending to a Match result can't overwrite the next rule's values.
func (g *MphMatcherGroup) valuesOf(ruleIdx uint32) []uint32 {
start, end := g.valueOffs[ruleIdx], g.valueOffs[ruleIdx+1]
return g.values[start:end:end]
}
// Lookup searches for input in minimal perfect hash table and returns its index. 0 indicates not found.
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 {
n := g.level1[i1]
// Build only puts valid rule indices in level1, so n+1 < len(patternOffs) and the span is inside patterns.
// Skip the bounds checks, they made this hot path measurably slower than indexing a []string
offs := (*[2]uint32)(unsafe.Add(unsafe.Pointer(unsafe.SliceData(g.patternOffs)), uintptr(n)*4))
if start := offs[0]; int(offs[1]-start) == len(input) && unsafe.String((*byte)(unsafe.Add(unsafe.Pointer(unsafe.StringData(g.patterns)), start)), len(input)) == input {
return n
}
return 0
@@ -160,12 +193,12 @@ func (g *MphMatcherGroup) Match(input string) []uint32 {
hash = hash*PrimeRK + uint32(input[i])
if input[i] == '.' {
if mphIdx := g.Lookup(hash, input[i:]); mphIdx != 0 {
matches = append(matches, g.values[mphIdx])
matches = append(matches, g.valuesOf(mphIdx))
}
}
}
if mphIdx := g.Lookup(hash, input); mphIdx != 0 {
matches = append(matches, g.values[mphIdx])
matches = append(matches, g.valuesOf(mphIdx))
}
return CompositeMatchesReverse(matches)
}
@@ -1,7 +1,9 @@
package strmatcher_test
import (
"math/rand"
"reflect"
"slices"
"testing"
"github.com/xtls/xray-core/common"
@@ -276,3 +278,63 @@ 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)
}
}
g.Build()
for _, input := range inputs {
keys := []string{input} // Whole input first, then "." suffixes from longest to shortest
for i := range len(input) {
if input[i] == '.' {
keys = append(keys, input[i:])
}
}
var want []uint32
for _, k := range keys {
want = append(append(want, full[k]...), domain[k]...)
}
if m := g.Match(input); !slices.Equal(m, want) {
t.Fatalf("seed %d: Match(%q) = %v, want %v", seed, input, m, want)
}
if m := g.MatchAny(input); m != (len(want) > 0) {
t.Fatalf("seed %d: MatchAny(%q) = %v", seed, input, m)
}
}
}
}
func TestMphMatcherGroupAppend(t *testing.T) {
g := NewMphMatcherGroup()
g.AddFullMatcher(FullMatcher("a.com"), 1)
g.AddFullMatcher(FullMatcher("b.com"), 2)
g.Build()
if m := append(g.Match("a.com"), 3); !slices.Equal(m, []uint32{1, 3}) {
t.Error("expect [1 3], but ", m)
}
if m := g.Match("b.com"); !slices.Equal(m, []uint32{2}) {
t.Error("expect [2], but ", m)
}
}
+45 -11
View File
@@ -3,6 +3,7 @@ package strmatcher
import (
"errors"
"regexp"
"regexp/syntax"
"slices"
"strings"
"unicode/utf8"
@@ -73,7 +74,43 @@ 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
}
func newRegexMatcher(pattern string) (Matcher, error) {
regex, err := regexp.Compile(pattern)
if err != nil {
return nil, err
}
m := &RegexMatcher{pattern: regex}
if re, err := syntax.Parse(pattern, syntax.Perl); err == nil { // same flags as regexp.Compile
m.literals = requiredLiterals(re, nil)
slices.SortStableFunc(m.literals, func(a, b string) int { return len(b) - len(a) })
}
return m, nil
}
// requiredLiterals appends to dst the case-sensitive strings that every match of re contains.
func requiredLiterals(re *syntax.Regexp, dst []string) []string {
switch re.Op {
case syntax.OpLiteral:
// regexp matches U+FFFD against invalid UTF-8 bytes, strings.Contains does not
if re.Flags&syntax.FoldCase == 0 && !slices.Contains(re.Rune, utf8.RuneError) {
dst = append(dst, string(re.Rune))
}
case syntax.OpCapture, syntax.OpPlus:
dst = requiredLiterals(re.Sub[0], dst)
case syntax.OpRepeat:
if re.Min > 0 {
dst = requiredLiterals(re.Sub[0], dst)
}
case syntax.OpConcat:
for _, sub := range re.Sub {
dst = requiredLiterals(sub, dst)
}
}
return dst
}
func (*RegexMatcher) Type() Type {
@@ -89,6 +126,11 @@ func (m *RegexMatcher) String() string {
}
func (m *RegexMatcher) Match(s string) bool {
for _, l := range m.literals {
if !strings.Contains(s, l) {
return false
}
}
return m.pattern.MatchString(s)
}
@@ -102,11 +144,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 +173,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")
}
@@ -0,0 +1,60 @@
package strmatcher
import (
"regexp"
"slices"
"testing"
)
var regexLiteralCases = []struct {
pattern string
literals []string
}{
{`(^|\.)91porn\.(best|com)$`, []string{"91porn."}},
{`.+\.awsdns-cn-[0-9][0-9]\.(biz|com|net|top)$`, []string{".awsdns-cn-", "."}},
{`^r+[0-9]+(---|\.)sn-(2x3|ni5|j5o)\w{5}\.googlevideo\.com$`, []string{".googlevideo.com", "sn-", "r"}},
{`(?i)abc`, nil},
{`ab(?i:CD)ef`, []string{"ab", "ef"}},
{`(abc)?x`, []string{"x"}},
{`(abc)*x`, []string{"x"}},
{`x{0,3}yy`, []string{"yy"}},
{`(ab)+c{2}`, []string{"ab", "c"}},
{`abc|abd`, []string{"ab"}},
{`\Qa.b\E`, []string{"a.b"}},
{`a\x{FFFD}b`, nil},
{`^[^.]+$`, nil},
}
func TestRegexRequiredLiterals(t *testing.T) {
for _, test := range regexLiteralCases {
m, err := newRegexMatcher(test.pattern)
if err != nil {
t.Fatal(err)
}
if got := m.(*RegexMatcher).literals; !slices.Equal(got, test.literals) {
t.Errorf("%s: got %q, want %q", test.pattern, got, test.literals)
}
}
}
func FuzzRegexMatcher(f *testing.F) {
inputs := []string{
"", "x", "yy", "abd", "ccc", "ABC", "abCDef", "abcdef", "abababcc", "a.b", "a\xffb", "a\uFFFDb",
"www.91porn.com", "ns1.awsdns-cn-01.top", "r1---sn-2x3abcde.googlevideo.com",
}
for _, test := range regexLiteralCases {
for _, s := range inputs {
f.Add(test.pattern, s)
}
}
f.Fuzz(func(t *testing.T, pattern, s string) {
re, err := regexp.Compile(pattern)
if err != nil {
return
}
m, _ := newRegexMatcher(pattern)
if got, want := m.Match(s), re.MatchString(s); got != want {
t.Errorf("pattern %q, input %q: got %v, want %v", pattern, s, got, want)
}
})
}
+21 -25
View File
@@ -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)
}
+4
View File
@@ -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()
+1 -1
View File
@@ -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 {
+1 -1
View File
@@ -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()
+1 -1
View File
@@ -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 {
+20
View File
@@ -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
}
+1 -1
View File
@@ -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()
-53
View File
@@ -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
}
-72
View File
@@ -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
}
-10
View File
@@ -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
}
-58
View File
@@ -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())
}
-70
View File
@@ -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) {
}
-107
View File
@@ -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
}
-81
View File
@@ -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
}
-66
View File
@@ -1,66 +0,0 @@
package singbridge
import (
"time"
"github.com/sagernet/sing/common"
"github.com/sagernet/sing/common/bufio"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
)
var (
_ buf.Reader = (*Conn)(nil)
_ buf.TimeoutReader = (*Conn)(nil)
_ buf.Writer = (*Conn)(nil)
)
type Conn struct {
net.Conn
writer N.VectorisedWriter
}
func NewConn(conn net.Conn) *Conn {
writer, _ := bufio.CreateVectorisedWriter(conn)
return &Conn{
Conn: conn,
writer: writer,
}
}
func (c *Conn) ReadMultiBuffer() (buf.MultiBuffer, error) {
buffer, err := buf.ReadBuffer(c.Conn)
if err != nil {
return nil, err
}
return buf.MultiBuffer{buffer}, nil
}
func (c *Conn) ReadMultiBufferTimeout(duration time.Duration) (buf.MultiBuffer, error) {
err := c.SetReadDeadline(time.Now().Add(duration))
if err != nil {
return nil, err
}
defer c.SetReadDeadline(time.Time{})
return c.ReadMultiBuffer()
}
func (c *Conn) WriteMultiBuffer(bufferList buf.MultiBuffer) error {
defer buf.ReleaseMulti(bufferList)
if c.writer != nil {
bytesList := make([][]byte, len(bufferList))
for i, buffer := range bufferList {
bytesList[i] = buffer.Bytes()
}
return common.Error(bufio.WriteVectorised(c.writer, bytesList))
}
// Since this conn is only used by tun, we don't force buffer writes to merge.
for _, buffer := range bufferList {
_, err := c.Conn.Write(buffer.Bytes())
if err != nil {
return err
}
}
return nil
}
+2 -2
View File
@@ -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)
}
+4 -4
View File
@@ -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) {
+1 -1
View File
@@ -20,7 +20,7 @@ import (
var (
Version_x byte = 26
Version_y byte = 9
Version_z byte = 8
Version_z byte = 9
)
var (
+1 -1
View File
@@ -13,7 +13,7 @@ type FakeDNSEngine interface {
var (
FakeIPv4Pool = "198.18.0.0/15"
FakeIPv6Pool = "fc00::/18"
FakeIPv6Pool = "2001:2::/48"
)
type FakeDNSEngineRev0 interface {
+10 -11
View File
@@ -18,25 +18,24 @@ 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
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba
golang.org/x/crypto v0.55.0
golang.org/x/crypto v0.57.0
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842
golang.org/x/net v0.58.0
golang.org/x/sync v0.22.0
golang.org/x/sys v0.47.0
golang.org/x/net v0.59.0
golang.org/x/sync v0.23.0
golang.org/x/sys v0.48.0
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
golang.zx2c4.com/wireguard/windows v1.0.1
google.golang.org/grpc v1.83.2
golang.zx2c4.com/wireguard/windows v1.1.1
google.golang.org/grpc v1.84.0
google.golang.org/protobuf v1.36.12
gvisor.dev/gvisor v0.0.0-20260122175437-89a5d21be8f0
h12.io/socks v1.0.3
lukechampine.com/blake3 v1.4.1
mvdan.cc/gofumpt v0.12.0
)
require (
@@ -48,7 +47,6 @@ require (
github.com/juju/ratelimit v1.0.2 // indirect
github.com/klauspost/compress v1.17.4 // indirect
github.com/koron/go-ssdp v0.0.4 // indirect
github.com/kr/text v0.2.0 // indirect
github.com/libp2p/go-netroute v0.2.1 // indirect
github.com/pion/dtls/v3 v3.1.5 // indirect
github.com/pion/logging v0.2.4 // indirect
@@ -57,8 +55,9 @@ require (
github.com/vishvananda/netns v0.0.5 // indirect
github.com/wlynxg/anet v0.0.5 // indirect
go.yaml.in/yaml/v3 v3.0.5 // indirect
golang.org/x/text v0.41.0 // indirect
golang.org/x/text v0.42.0 // indirect
golang.org/x/time v0.14.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect
golang.org/x/tools v0.49.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 // indirect
gopkg.in/yaml.v2 v2.4.0 // indirect
)
+24 -41
View File
@@ -2,17 +2,12 @@ github.com/andybalholm/brotli v1.0.6 h1:Yf9fFpf49Zrxb9NlQaluyE92/+X7UVHlhMNJN2sx
github.com/andybalholm/brotli v1.0.6/go.mod h1:fO7iG3H7G2nSZ7m0zPUDn85XEX2GTukHGRSepvi9Eig=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e h1:5mgtR5gwIgBKMiGI1QdXldZZ+SNor06Nbu1wCBulQBg=
github.com/apernet/quic-go v0.61.1-0.20260806010916-184d081eef3e/go.mod h1:x7qxEvX6MCVtDuBKHj3E+88+BtrbEMuAL5qGUKItjW8=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/cloudflare/circl v1.6.5 h1:O64F26HEqNhznd/hrC5KZXVKYuKM2rx4deZDTc4ihQA=
github.com/cloudflare/circl v1.6.5/go.mod h1:h5LNyxAc5nTue9DS5jT+48en2PSDYt3zdGnz5OstK6c=
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344 h1:Arcl6UOIS/kgO2nW3A65HN+7CMjSDP/gofXL4CZt1V4=
github.com/ghodss/yaml v1.0.1-0.20220118164431-d8423dcdf344/go.mod h1:GIjDIg/heH5DOkXY3YJ/wNhfHsQHoXGjl8G8amsYQ1I=
github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
github.com/go-quicktest/qt v1.102.0 h1:HSQxCeh5YZH3EL3W39ixjtyaEhcWSXQHtHnMBzSs474=
github.com/go-quicktest/qt v1.102.0/go.mod h1:p4lGIVX+8Wa6ZPNDvqcxq36XpUDLh42FLetFU7odllI=
github.com/golang/mock v1.7.0-rc.1 h1:YojYx61/OLFsiv6Rw1Z96LpldJIy31o+UHmwAUMJ6/U=
github.com/golang/mock v1.7.0-rc.1/go.mod h1:s42URUywIqd+OcERslBJvOjepvNymP31m3q8d/GkuRs=
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
@@ -73,12 +68,8 @@ github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af h1:er
github.com/refraction-networking/utls v1.8.3-0.20260301010127-aa6edf4b11af/go.mod h1:jkSOEkLqn+S/jtpEHPOsVv/4V4EVnelwbMQl4vCWXAM=
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.10.0 h1:TMyTOH3F/DB16zRVcYyreMH6GnZZrwQVAoYjRBZyWFQ=
github.com/rogpeppe/go-internal v1.10.0/go.mod h1:UQnix2H7Ngw/k4C5ijL5+65zddjncjaFoBhdsK/akog=
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/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/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=
@@ -90,18 +81,6 @@ github.com/wlynxg/anet v0.0.5/go.mod h1:eay5PRQr7fIVAMbTbchTnO9gG65Hg/uYGdc7mguH
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0 h1:rb+fKQFhz+5I2PPuQsNYxI5mUU840XWYtRF0ZBjvkws=
github.com/xtls/reality v0.0.0-20260908062103-8cdf7bf9c7f0/go.mod h1:DsJblcWDGt76+FVqBVwbwRhxyyNJsGV48gJLch0OOWI=
github.com/yuin/goldmark v1.4.1/go.mod h1:mwnBkeHKe2W/ZEtQ+71ViKU8L12m81fl3OWwC1Zlc8k=
go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
go.opentelemetry.io/otel v1.44.0 h1:JjwHmHpA4iZ3wBxluu2fbbE7j4kqlE8jXyAyPXH7HqU=
go.opentelemetry.io/otel v1.44.0/go.mod h1:BMgjTHL9WPRlRjL2oZCBTL4whCGtXch2H4BhOPIAyYc=
go.opentelemetry.io/otel/metric v1.44.0 h1:1w0gILTcHdr3YI+ixLyjemwrVnsMURbTZFrSYCdDdmc=
go.opentelemetry.io/otel/metric v1.44.0/go.mod h1:8O7hanEPBNgEMmybD3s2VBKcgWOCsA6tzHBPODAiquo=
go.opentelemetry.io/otel/sdk v1.44.0 h1:nHYwb9lK+fJPU/dnT6s7W7Z8itMWyqrnVfbheVYrZ58=
go.opentelemetry.io/otel/sdk v1.44.0/go.mod h1:Osuydd3Se74nqjAKxid74N5eC+jfEqfTegHRnq58oK0=
go.opentelemetry.io/otel/sdk/metric v1.44.0 h1:3LlKgI+VjbVsjNRFZJZAJ30WjXC5VkNRks6si09iEfI=
go.opentelemetry.io/otel/sdk/metric v1.44.0/go.mod h1:5B5pMARnXxKhltooO4xUuCBorl65a4EpnTalObqOigA=
go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/6gtIk=
go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE=
go.uber.org/mock v0.5.2 h1:LbtPTcP8A5k9WPXj54PPPbjcI4Y6lhyOZXn+VS7wNko=
go.uber.org/mock v0.5.2/go.mod h1:wLlUxC2vVTPTaE3UD51E0BGOAElKrILxhVSDYQLld5o=
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
@@ -110,8 +89,8 @@ go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBs
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI=
golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M=
golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis=
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842 h1:vr/HnozRka3pE4EsMEg1lgkXJkTFJCVUX+S/ZT6wYzM=
golang.org/x/exp v0.0.0-20240506185415-9bf2ced13842/go.mod h1:XtvwrStGgqGPLc4cjQfWqZHG1YFdYs6swckp8vpsjnc=
golang.org/x/lint v0.0.0-20200302205851-738671d3881b/go.mod h1:3xt1FjdF8hUf6vQPIChWIBhFzV8gjjsPE/fR3IyQdNY=
@@ -120,12 +99,12 @@ golang.org/x/mod v0.5.1/go.mod h1:5OXOZSfqPIIbmVBIIKWRFfZjPR0E5r58TLhUjH0a2Ro=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s=
golang.org/x/net v0.0.0-20211015210444-4f30a5c0130f/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y=
golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To=
golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU=
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues=
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20210220032951-036812b2e83c/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek=
golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sync v0.23.0 h1:KameEIfc1IkluZyXWLn39Wd4tURc6GbCiISGiZm2bQk=
golang.org/x/sync v0.23.0/go.mod h1:sUUOizhqBxiL6pEWpqNLUiaJn1ShEbZ6BBqskPbjZm0=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
@@ -133,20 +112,22 @@ golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7w
golang.org/x/sys v0.0.0-20211019181941-9d821ace8654/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs=
golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ=
golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8=
golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M=
golang.org/x/text v0.42.0 h1:JbOZXgfeCPU9gacVtYliJqOhD+zhrEqK4LfdpmlUZqI=
golang.org/x/text v0.42.0/go.mod h1:ojzP1Z+2QtioaF8DTtO8K5q7JWVVYwZKenzujK0Zd0E=
golang.org/x/time v0.14.0 h1:MRx4UaLrDotUKUdCIqzPC48t1Y9hANFKIRpNx+Te8PI=
golang.org/x/time v0.14.0/go.mod h1:eL/Oa2bBBK0TkX57Fyni+NgnyQQN4LitPmob2Hjnqw4=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20191119224855-298f0cb1881e/go.mod h1:b+2E5dAYhXwXZwtnZ6UAqBI28+e2cm9otk0dWdXHAEo=
golang.org/x/tools v0.0.0-20200130002326-2f3ba24bd6e7/go.mod h1:TB2adYChydJhpapKDTa4BR/hXlZSLoq2Wpct/0txZ28=
golang.org/x/tools v0.1.8/go.mod h1:nABZi5QlRsZVlzPpHl034qft6wpY4eDcsTt5AaioBiU=
golang.org/x/tools v0.49.0 h1:3NI7VXzL9+1WZD52Dx2ttoPwD5DWrFGpl9mFZDlmisI=
golang.org/x/tools v0.49.0/go.mod h1:SJNXV9DBKT0UbdttsQjbfJlAE/q+y36++zo3uL3N0Oo=
golang.org/x/xerrors v0.0.0-20190717185122-a985d3407aa7/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20191011141410-1b5146add898/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
@@ -154,14 +135,14 @@ golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeu
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
golang.zx2c4.com/wireguard/windows v1.1.1 h1:8/H97U1v1PNDNcBsMZgU3KFuND9MQdTsU2NOwmCXArE=
golang.zx2c4.com/wireguard/windows v1.1.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4=
gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa h1:mZHHdPZl0dbGHCflZgAq/Q468DWVFcU2whhB2KAo8fk=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU=
google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800 h1:qEHAMpSaUhtD0p3NbEEI83HwNGFxEwaSJ1G9PLnCBZE=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260706201446-f0a921348800/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8=
google.golang.org/grpc v1.84.0 h1:soMyaPJ8pAak5PIQ0DGBUir0XRo2fRoMqhNWMLlLxO0=
google.golang.org/grpc v1.84.0/go.mod h1:ljCht0DrxQrXBDRTZp52Qxh3Ffk8CdYm2sj4O2QN2C0=
google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc=
google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
@@ -176,3 +157,5 @@ h12.io/socks v1.0.3 h1:Ka3qaQewws4j4/eDQnOdpr4wXsC//dXtWvftlIcCQUo=
h12.io/socks v1.0.3/go.mod h1:AIhxy1jOId/XCz9BO+EIgNL2rQiPTBNnOfnVnQ+3Eck=
lukechampine.com/blake3 v1.4.1 h1:I3Smz7gso8w4/TunLKec6K2fn+kyKtDxr/xcQEN84Wg=
lukechampine.com/blake3 v1.4.1/go.mod h1:QFosUxmjB8mnrWFSNwKmvxHpfY72bmD2tQ0kBMM3kwo=
mvdan.cc/gofumpt v0.12.0 h1:1Lbudkz2kpM9Cjz2pL4M19u7q+GaEhCTNf7N9mfpcho=
mvdan.cc/gofumpt v0.12.0/go.mod h1:SmBHHrljiZu/uoypeKup3rFzP6eoC9UwCp2iH5E3jZA=
+2 -2
View File
@@ -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())
+1 -1
View File
@@ -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
+2 -2
View File
@@ -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 {
+103
View File
@@ -0,0 +1,103 @@
package conf
import (
"net/netip"
"strings"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/proxy/masque"
"google.golang.org/protobuf/proto"
)
type MasqueClientConfig struct {
Address *Address `json:"address"`
Port uint16 `json:"port"`
RemoteDNS []string `json:"remoteDNS"`
}
func (c *MasqueClientConfig) Build() (proto.Message, error) {
if c.Address == nil {
return nil, errors.New(`MASQUE: "address" is not set`)
}
if c.Port == 0 {
return nil, errors.New(`MASQUE: "port" is not set`)
}
for _, s := range c.RemoteDNS {
if _, err := netip.ParseAddr(s); err != nil {
return nil, errors.New(`MASQUE: invalid "remoteDNS" `, s).Base(err)
}
}
return &masque.ClientConfig{
Server: &protocol.ServerEndpoint{
Address: c.Address.Build(),
Port: uint32(c.Port),
},
RemoteDns: c.RemoteDNS,
}, nil
}
type MasqueUserConfig struct {
Pass string `json:"pass"`
Level uint32 `json:"level"`
Email string `json:"email"`
}
type MasqueServerConfig struct {
Users []*MasqueUserConfig `json:"users"`
Clients []*MasqueUserConfig `json:"clients"`
Address []string `json:"address"`
MTU uint32 `json:"mtu"`
}
func (c *MasqueServerConfig) Build() (proto.Message, error) {
if c.Clients != nil {
c.Users = c.Clients
}
config := &masque.ServerConfig{
Address: c.Address,
Mtu: c.MTU,
}
emails := make(map[string]bool)
for _, user := range c.Users {
if user.Email == "" {
return nil, errors.New(`MASQUE: "email" is empty`)
}
if strings.Contains(user.Email, ":") {
return nil, errors.New(`MASQUE: invalid "email" `, user.Email)
}
if user.Pass == "" {
return nil, errors.New(`MASQUE: "pass" of `, user.Email, ` is empty`)
}
email := strings.ToLower(user.Email)
if emails[email] {
return nil, errors.New(`MASQUE: duplicate "email" `, user.Email)
}
emails[email] = true
config.Users = append(config.Users, &protocol.User{
Email: user.Email,
Level: user.Level,
Account: serial.ToTypedMessage(&masque.Account{Password: user.Pass}),
})
}
if len(c.Address) == 0 {
return nil, errors.New(`MASQUE: "address" is not set`)
}
var v4, v6 bool
for _, s := range c.Address {
prefix, err := netip.ParsePrefix(s)
if err != nil {
return nil, errors.New(`MASQUE: invalid "address" `, s).Base(err)
}
if prefix.Addr().Is4() && v4 || prefix.Addr().Is6() && v6 {
return nil, errors.New(`MASQUE: "address" takes at most one IPv4 and one IPv6 prefix`)
}
v4 = v4 || prefix.Addr().Is4()
v6 = v6 || prefix.Addr().Is6()
}
if c.MTU != 0 && (c.MTU < 1280 || c.MTU > 65535) {
return nil, errors.New(`MASQUE: "mtu" must be between 1280 and 65535`)
}
return config, nil
}
+188
View File
@@ -0,0 +1,188 @@
package conf_test
import (
"encoding/json"
"testing"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
. "github.com/xtls/xray-core/infra/conf"
masqueproxy "github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/transport/internet/masque"
)
func TestMasqueConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/.well-known/masque/ip/*/*/"},
},
{
Input: `{
"host": "example.com:8443",
"path": "/.well-known/masque/ip/{target}/{ipproto}/",
"headers": {"Authorization": "Basic dTpw"}
}`,
Parser: loadJSON(creator),
Output: &masque.Config{
Host: "example.com:8443",
Path: "/.well-known/masque/ip/*/*/",
Headers: map[string]string{"Authorization": "Basic dTpw"},
},
},
{
Input: `{"path": "/masque/ip{?target,ipproto}"}`,
Parser: loadJSON(creator),
Output: &masque.Config{Path: "/masque/ip?target=*&ipproto=*"},
},
{
Input: `{"user": "u", "pass": "p:q", "headers": {"X-Token": "a"}}`,
Parser: loadJSON(creator),
Output: &masque.Config{
Path: "/.well-known/masque/ip/*/*/",
Headers: map[string]string{"Authorization": "Basic dTpwOnE=", "X-Token": "a"},
},
},
})
for _, input := range []string{
`{"path": "/masque/{target}/{ipproto}/{dns}"}`,
`{"path": "masque"}`,
`{"host": "example.com/path"}`,
`{"headers": {"host": "example.com"}}`,
`{"headers": {"Capsule-Protocol": "?0"}}`,
`{"headers": {"X Token": "a"}}`,
`{"headers": {"X-Token": "a\r\nb"}}`,
`{"user": "u:v", "pass": "p"}`,
`{"user": "u", "pass": "p", "headers": {"authorization": "Basic dTpw"}}`,
} {
if _, err := loadJSON(creator)(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueOutboundConfig(t *testing.T) {
build := func(s string) error {
c := new(OutboundDetourConfig)
if err := json.Unmarshal([]byte(s), c); err != nil {
return err
}
_, err := c.Build()
return err
}
if err := build(`{
"protocol": "masque",
"settings": {"address": "example.com", "port": 443},
"streamSettings": {"network": "masque", "security": "tls"},
"mux": {"enabled": false, "concurrency": -1}
}`); err != nil {
t.Error(err)
}
for _, input := range []string{
`{"protocol": "masque", "settings": {"address": "example.com"}, "streamSettings": {"network": "masque", "security": "tls"}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true}}`,
`{"protocol": "masque", "settings": {"address": "example.com", "port": 443}, "streamSettings": {"network": "masque", "security": "tls"}, "mux": {"enabled": true, "concurrency": -1}}`,
`{"protocol": "freedom", "streamSettings": {"network": "masque", "security": "tls"}}`,
} {
if err := build(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueServerConfig(t *testing.T) {
creator := func() Buildable {
return new(MasqueServerConfig)
}
runMultiTestCase(t, []TestCase{
{
Input: `{
"users": [{"email": "u@example.com", "pass": "p", "level": 1}],
"address": ["10.13.0.1/24", "fd13::1/64"],
"mtu": 1400
}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Users: []*protocol.User{{
Email: "u@example.com",
Level: 1,
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p"}),
}},
Address: []string{"10.13.0.1/24", "fd13::1/64"},
Mtu: 1400,
},
},
{
Input: `{"clients": [{"email": "u", "pass": "p:q"}], "address": ["10.13.0.1/24"]}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Users: []*protocol.User{{
Email: "u",
Account: serial.ToTypedMessage(&masqueproxy.Account{Password: "p:q"}),
}},
Address: []string{"10.13.0.1/24"},
},
},
{
Input: `{"address": ["10.13.0.1/24"]}`,
Parser: loadJSON(creator),
Output: &masqueproxy.ServerConfig{
Address: []string{"10.13.0.1/24"},
},
},
})
for _, input := range []string{
`{"users": [{"email": "u:v", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "", "pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"pass": "p"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": ""}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}, {"email": "U", "pass": "q"}], "address": ["10.13.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24", "10.14.0.1/24"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["fd13::1/64", "fd14::1/64"]}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 1000}`,
`{"users": [{"email": "u", "pass": "p"}], "address": ["10.13.0.1/24"], "mtu": 70000}`,
} {
if _, err := loadJSON(creator)(input); err == nil {
t.Errorf("expected an error for %s", input)
}
}
}
func TestMasqueInboundConfig(t *testing.T) {
build := func(s string) error {
c := new(InboundDetourConfig)
if err := json.Unmarshal([]byte(s), c); err != nil {
return err
}
_, err := c.Build()
return err
}
if err := build(`{
"protocol": "masque",
"port": 443,
"settings": {"users": [{"email": "u@example.com", "pass": "p"}], "address": ["10.13.0.1/24"]},
"streamSettings": {"network": "masque", "security": "tls"}
}`); err != nil {
t.Error(err)
}
if err := build(`{
"protocol": "vless",
"port": 443,
"settings": {"users": [{"id": "27848739-7e62-4138-9fd3-098a63964b6b"}], "decryption": "none"},
"streamSettings": {"network": "masque", "security": "tls"}
}`); err == nil {
t.Error("expected an error for the masque transport on a vless inbound")
}
}
+1 -1
View File
@@ -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) {
+37 -56
View File
@@ -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
}
+2 -3
View File
@@ -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())
+55 -1
View File
@@ -23,6 +23,7 @@ import (
"github.com/xtls/xray-core/transport/internet/finalmask/realm"
"github.com/xtls/xray-core/transport/internet/finalmask/salamander"
"github.com/xtls/xray-core/transport/internet/finalmask/sudoku"
"github.com/xtls/xray-core/transport/internet/finalmask/udphop"
"github.com/xtls/xray-core/transport/internet/finalmask/xdns"
"github.com/xtls/xray-core/transport/internet/finalmask/xicmp"
"github.com/xtls/xray-core/transport/internet/finalmask/xmc"
@@ -83,6 +84,7 @@ var (
"xdns": func() interface{} { return new(Xdns) },
"xicmp": func() interface{} { return new(Xicmp) },
"realm": func() interface{} { return new(Realm) },
"udphop": func() interface{} { return new(UDPHop) },
}, "type", "settings")
)
@@ -905,6 +907,59 @@ func (c *Realm) Build() (proto.Message, error) {
}, nil
}
type UDPHop struct {
Mode string `json:"mode"`
Interval Int32Range `json:"interval"`
RemoteIPs []string `json:"remoteIPs"`
RemotePorts PortList `json:"remotePorts"`
}
func (c *UDPHop) Build() (proto.Message, error) {
var local, remote, remoteOnce bool
for _, mode := range strings.Split(c.Mode, ",") {
switch strings.ToLower(mode) {
case "intervallocal":
local = true
case "intervalremote":
remote = true
case "perconnremote":
remoteOnce = true
default:
return nil, errors.New("invalid mode ", mode)
}
}
var remoteIPs []string
for _, ip := range c.RemoteIPs {
prefix, err := netip.ParsePrefix(ip)
if err == nil {
remoteIPs = append(remoteIPs, prefix.String())
continue
}
addr, err := netip.ParseAddr(ip)
if err == nil {
remoteIPs = append(remoteIPs, netip.PrefixFrom(addr, addr.BitLen()).String())
continue
}
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(interval.From),
IntervalMax: int64(interval.To),
RemoteIPs: remoteIPs,
RemotePorts: c.RemotePorts.Build().Ports(),
}, nil
}
type Mask struct {
Type string `json:"type"`
Settings *json.RawMessage `json:"settings"`
@@ -938,7 +993,6 @@ type QuicParamsConfig struct {
BrutalUp Bandwidth `json:"brutalUp"`
BrutalDown Bandwidth `json:"brutalDown"`
BrutalDisableLossCompensation bool `json:"brutalDisableLossCompensation"`
UdpHop UdpHop `json:"udpHop"`
InitStreamReceiveWindow uint64 `json:"initStreamReceiveWindow"`
MaxStreamReceiveWindow uint64 `json:"maxStreamReceiveWindow"`
InitConnectionReceiveWindow uint64 `json:"initConnectionReceiveWindow"`
+37 -20
View File
@@ -36,6 +36,10 @@ func (p TransportProtocol) Build() (string, error) {
return "", errors.PrintRemovedFeatureError("QUIC transport (without web service, etc.)", "XHTTP stream-one H3")
case "hysteria":
return "hysteria", nil
case "masque":
return "masque", nil
case "xdrive":
return "xdrive", nil
default:
return "", errors.New("Config: unknown transport protocol: ", p)
}
@@ -59,6 +63,8 @@ type StreamConfig struct {
WSSettings *WebSocketConfig `json:"wsSettings"`
HTTPUPGRADESettings *HttpUpgradeConfig `json:"httpupgradeSettings"`
HysteriaSettings *HysteriaConfig `json:"hysteriaSettings"`
MASQUESettings *MasqueConfig `json:"masqueSettings"`
XDRIVESettings *XDriveConfig `json:"xdriveSettings"`
SocketSettings *SocketConfig `json:"sockopt"`
}
@@ -192,6 +198,26 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
Settings: serial.ToTypedMessage(hs),
})
}
if c.MASQUESettings != nil {
ms, err := c.MASQUESettings.Build()
if err != nil {
return nil, errors.New("Failed to build MASQUE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "masque",
Settings: serial.ToTypedMessage(ms),
})
}
if c.XDRIVESettings != nil {
xs, err := c.XDRIVESettings.Build()
if err != nil {
return nil, errors.New("Failed to build XDRIVE config.").Base(err)
}
config.TransportSettings = append(config.TransportSettings, &internet.TransportConfig{
ProtocolName: "xdrive",
Settings: serial.ToTypedMessage(xs),
})
}
if c.SocketSettings != nil {
ss, err := c.SocketSettings.Build()
if err != nil {
@@ -253,10 +279,6 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
return nil, errors.New("unknown congestion control: ", c.FinalMask.QuicParams.Congestion, ", valid values: reno, bbr, brutal, force-brutal")
}
if (c.FinalMask.QuicParams.UdpHop.Interval.From != 0 && c.FinalMask.QuicParams.UdpHop.Interval.From < 5) || (c.FinalMask.QuicParams.UdpHop.Interval.To != 0 && c.FinalMask.QuicParams.UdpHop.Interval.To < 5) {
return nil, errors.New("Interval must be at least 5")
}
if c.FinalMask.QuicParams.InitStreamReceiveWindow > 0 && c.FinalMask.QuicParams.InitStreamReceiveWindow < 16384 {
return nil, errors.New("InitStreamReceiveWindow must be at least 16384")
}
@@ -290,22 +312,17 @@ func (c *StreamConfig) Build() (*internet.StreamConfig, error) {
BrutalUp: up,
BrutalDown: down,
BrutalDisableLossCompensation: c.FinalMask.QuicParams.BrutalDisableLossCompensation,
UdpHop: &internet.UdpHop{
Ports: c.FinalMask.QuicParams.UdpHop.PortList.Build().Ports(),
IntervalMin: int64(c.FinalMask.QuicParams.UdpHop.Interval.From),
IntervalMax: int64(c.FinalMask.QuicParams.UdpHop.Interval.To),
},
InitStreamReceiveWindow: c.FinalMask.QuicParams.InitStreamReceiveWindow,
MaxStreamReceiveWindow: c.FinalMask.QuicParams.MaxStreamReceiveWindow,
InitConnReceiveWindow: c.FinalMask.QuicParams.InitConnectionReceiveWindow,
MaxConnReceiveWindow: c.FinalMask.QuicParams.MaxConnectionReceiveWindow,
MaxIdleTimeout: c.FinalMask.QuicParams.MaxIdleTimeout,
KeepAlivePeriod: c.FinalMask.QuicParams.KeepAlivePeriod,
DisablePathMtuDiscovery: c.FinalMask.QuicParams.DisablePathMTUDiscovery,
DisableChromeParrot: c.FinalMask.QuicParams.DisableChromeParrot,
DisableGSO: c.FinalMask.QuicParams.DisableGSO,
MaxIncomingStreams: c.FinalMask.QuicParams.MaxIncomingStreams,
DisableStatelessReset: c.FinalMask.QuicParams.DisableStatelessReset,
InitStreamReceiveWindow: c.FinalMask.QuicParams.InitStreamReceiveWindow,
MaxStreamReceiveWindow: c.FinalMask.QuicParams.MaxStreamReceiveWindow,
InitConnReceiveWindow: c.FinalMask.QuicParams.InitConnectionReceiveWindow,
MaxConnReceiveWindow: c.FinalMask.QuicParams.MaxConnectionReceiveWindow,
MaxIdleTimeout: c.FinalMask.QuicParams.MaxIdleTimeout,
KeepAlivePeriod: c.FinalMask.QuicParams.KeepAlivePeriod,
DisablePathMtuDiscovery: c.FinalMask.QuicParams.DisablePathMTUDiscovery,
DisableChromeParrot: c.FinalMask.QuicParams.DisableChromeParrot,
DisableGSO: c.FinalMask.QuicParams.DisableGSO,
MaxIncomingStreams: c.FinalMask.QuicParams.MaxIncomingStreams,
DisableStatelessReset: c.FinalMask.QuicParams.DisableStatelessReset,
}
}
}
+119 -30
View File
@@ -1,8 +1,9 @@
package conf
import (
"context"
"encoding/base64"
"encoding/json"
"maps"
"math/big"
"net/url"
"sort"
@@ -21,9 +22,12 @@ import (
"github.com/xtls/xray-core/transport/internet/httpupgrade"
"github.com/xtls/xray-core/transport/internet/hysteria"
"github.com/xtls/xray-core/transport/internet/kcp"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/splithttp"
"github.com/xtls/xray-core/transport/internet/tcp"
"github.com/xtls/xray-core/transport/internet/websocket"
"github.com/xtls/xray-core/transport/internet/xdrive"
"golang.org/x/net/http/httpguts"
"google.golang.org/protobuf/proto"
)
@@ -122,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,
@@ -190,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,
@@ -240,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)
}
@@ -534,10 +538,6 @@ type KCPConfig struct {
// Build implements Buildable.
func (c *KCPConfig) Build() (proto.Message, error) {
if c.HeaderConfig != nil || c.Seed != nil {
return nil, errors.PrintRemovedFeatureError("mkcp header & seed", "finalmask/udp header-* & mkcp-original & mkcp-aes128gcm")
}
config := common.Must2(internet.CreateTransportConfig(kcp.ProtocolName)).(*kcp.Config)
if c.Mtu != nil {
@@ -560,16 +560,16 @@ func (c *KCPConfig) Build() (proto.Message, error) {
}
if config.Mtu < 21 {
return nil, errors.New("Mtu must be at least 21").AtError()
return nil, errors.New("MTU must be at least 21")
}
if config.Tti < 10 || config.Tti > 1000 {
return nil, errors.New("invalid mKCP TTI: ", c.Tti).AtError()
return nil, errors.New("TTI must be between 10 and 1000")
}
if config.CwndMultiplier < 1 {
return nil, errors.New("CwndMultiplier must be at least 1").AtError()
return nil, errors.New("CwndMultiplier must be at least 1")
}
if config.GetSendingBufferSize() == 0 {
return nil, errors.New("MaxSendingWindow must be >= Mtu").AtError()
return nil, errors.New("MaxSendingWindow must be at least ", config.Mtu)
}
return config, nil
@@ -739,11 +739,6 @@ func (b Bandwidth) Bps() (uint64, error) {
return uint64(val*float64(mul)) / 8, nil
}
type UdpHop struct {
PortList PortList `json:"ports"`
Interval Int32Range `json:"interval"`
}
type Masquerade struct {
Type string `json:"type"`
@@ -760,14 +755,8 @@ type Masquerade struct {
}
type HysteriaConfig struct {
Version int32 `json:"version"`
Auth string `json:"auth"`
Congestion *string `json:"congestion"`
Up *Bandwidth `json:"up"`
Down *Bandwidth `json:"down"`
UdpHop *UdpHop `json:"udphop"`
Version int32 `json:"version"`
Auth string `json:"auth"`
UdpIdleTimeout int64 `json:"udpIdleTimeout"`
Masquerade Masquerade `json:"masquerade"`
}
@@ -777,10 +766,6 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return nil, errors.New("version != 2")
}
if c.Congestion != nil || c.Up != nil || c.Down != nil || c.UdpHop != nil {
errors.LogWarning(context.Background(), "congestion & up & down & udphop move to finalmask/quicParams")
}
if c.UdpIdleTimeout != 0 && (c.UdpIdleTimeout < 2 || c.UdpIdleTimeout > 600) {
return nil, errors.New("UdpIdleTimeout must be between 2 and 600")
}
@@ -805,6 +790,63 @@ func (c *HysteriaConfig) Build() (proto.Message, error) {
return config, nil
}
type MasqueConfig struct {
Host string `json:"host"`
Path string `json:"path"`
User string `json:"user"`
Pass string `json:"pass"`
Headers map[string]string `json:"headers"`
}
func (c *MasqueConfig) Build() (proto.Message, error) {
path := c.Path
if path == "" {
path = masque.DefaultPath
}
path = strings.NewReplacer(
"{target}", "*", "{ipproto}", "*",
"{?target,ipproto}", "?target=*&ipproto=*", "{?ipproto,target}", "?ipproto=*&target=*",
"{&target,ipproto}", "&target=*&ipproto=*", "{&ipproto,target}", "&ipproto=*&target=*",
).Replace(path)
if !strings.HasPrefix(path, "/") || strings.ContainsAny(path, "{}") {
return nil, errors.New(`invalid "path": `, path, `, only the variables {target} and {ipproto} are supported`)
}
if c.Host != "" {
if u, err := url.Parse("https://" + c.Host); err != nil || u.Host != c.Host {
return nil, errors.New(`invalid "host": `, c.Host)
}
}
for k, v := range c.Headers {
if !httpguts.ValidHeaderFieldName(k) || !httpguts.ValidHeaderFieldValue(v) {
return nil, errors.New(`invalid header in "headers": `, strconv.Quote(k))
}
switch strings.ToLower(k) {
case "host", "capsule-protocol":
return nil, errors.New(`"headers" can't contain "`, k, `"`)
case "authorization":
if c.User != "" || c.Pass != "" {
return nil, errors.New(`"headers" can't contain "`, k, `" when "user" or "pass" is set`)
}
}
}
headers := c.Headers
if c.User != "" || c.Pass != "" {
if strings.Contains(c.User, ":") {
return nil, errors.New(`invalid "user": `, c.User)
}
headers = maps.Clone(c.Headers)
if headers == nil {
headers = make(map[string]string)
}
headers["Authorization"] = "Basic " + base64.StdEncoding.EncodeToString([]byte(c.User+":"+c.Pass))
}
return &masque.Config{
Host: c.Host,
Path: path,
Headers: headers,
}, nil
}
func readFileOrString(f string, s []string) ([]byte, error) {
if len(f) > 0 {
return filesystem.ReadCert(f)
@@ -814,3 +856,50 @@ func readFileOrString(f string, s []string) ([]byte, error) {
}
return nil, errors.New("both file and bytes are empty.")
}
type XDriveConfig struct {
RemoteFolder string `json:"remoteFolder"`
Service string `json:"service"`
Secrets []string `json:"secrets"`
SegmentBytes uint32 `json:"segmentBytes"`
FlushIntervalMs uint32 `json:"flushIntervalMs"`
PollIntervalMs uint32 `json:"pollIntervalMs"`
MaxPollIntervalMs uint32 `json:"maxPollIntervalMs"`
SessionTTLSeconds uint32 `json:"sessionTtlSeconds"`
Concurrency uint32 `json:"concurrency"`
EagerWindowMs uint32 `json:"eagerWindowMs"`
HoleTimeoutMs uint32 `json:"holeTimeoutMs"`
Template json.RawMessage `json:"template"`
}
// Build implements Buildable.
func (c *XDriveConfig) Build() (proto.Message, error) {
switch c.Service {
case "local":
case "Google Drive":
if len(c.Secrets) != 3 {
return nil, errors.New("Google Drive needs 3 secrets in order of ClientID, ClientSecret, RefreshToken")
}
case "template":
if len(c.Template) == 0 {
return nil, errors.New(`service "template" needs a "template" object`)
}
default:
return nil, errors.New("unsupported service")
}
config := &xdrive.Config{
RemoteFolder: c.RemoteFolder,
Service: c.Service,
Secrets: c.Secrets,
SegmentBytes: c.SegmentBytes,
FlushIntervalMs: c.FlushIntervalMs,
PollIntervalMs: c.PollIntervalMs,
MaxPollIntervalMs: c.MaxPollIntervalMs,
SessionTtlSeconds: c.SessionTTLSeconds,
Concurrency: c.Concurrency,
EagerWindowMs: c.EagerWindowMs,
HoleTimeoutMs: c.HoleTimeoutMs,
Template: string(c.Template),
}
return config, nil
}
+73
View File
@@ -291,3 +291,76 @@ func TestHeaderCustomUDPBuildRejectsExprWithoutArgs(t *testing.T) {
t.Fatalf("expected transform arg rejection, got %v", err)
}
}
func TestXDriveStreamConfig(t *testing.T) {
config := new(StreamConfig)
if err := json.Unmarshal([]byte(`{
"method": "xdrive",
"xdriveSettings": {
"remoteFolder": "/tmp/xdrive",
"service": "local"
}
}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
built, err := config.Build()
if err != nil {
t.Fatalf("Build: %v", err)
}
if built.ProtocolName != "xdrive" {
t.Fatalf("ProtocolName is %q, want %q", built.ProtocolName, "xdrive")
}
if len(built.TransportSettings) != 1 || built.TransportSettings[0].ProtocolName != "xdrive" {
t.Fatalf("TransportSettings is %v, want a single xdrive entry", built.TransportSettings)
}
}
func TestXDriveRejectsUnknownService(t *testing.T) {
config := new(XDriveConfig)
if err := json.Unmarshal([]byte(`{"remoteFolder": "/tmp/xdrive", "service": "Dropbox"}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
if _, err := config.Build(); err == nil {
t.Fatal("Build accepted an unsupported service")
}
}
func TestXDriveTemplateStreamConfig(t *testing.T) {
config := new(StreamConfig)
if err := json.Unmarshal([]byte(`{
"method": "xdrive",
"xdriveSettings": {
"remoteFolder": "folder",
"service": "template",
"secrets": ["user", "pass"],
"template": {
"flatten": true,
"auth": {"type": "basic", "username": "{secret0}", "password": "{secret1}"},
"put": {"method": "PUT", "url": "https://dav.example/{folder}/{name}"},
"get": {"method": "GET", "url": "https://dav.example/{folder}/{name}"},
"delete": {"method": "DELETE", "url": "https://dav.example/{folder}/{name}"},
"list": {"method": "PROPFIND", "url": "https://dav.example/{folder}/", "namesRegex": "<d:href>/folder/([^<]+)</d:href>"}
}
}
}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
built, err := config.Build()
if err != nil {
t.Fatalf("Build: %v", err)
}
if built.ProtocolName != "xdrive" {
t.Fatalf("ProtocolName is %q, want xdrive", built.ProtocolName)
}
}
func TestXDriveTemplateNeedsTemplate(t *testing.T) {
config := new(XDriveConfig)
if err := json.Unmarshal([]byte(`{"remoteFolder": "f", "service": "template"}`), config); err != nil {
t.Fatalf("Unmarshal: %v", err)
}
if _, err := config.Build(); err == nil {
t.Fatal("Build accepted a template service without a template")
}
}
+2
View File
@@ -20,6 +20,7 @@ type TunConfig struct {
UserLevel uint32 `json:"userLevel"`
AutoSystemRoutingTable []string `json:"autoSystemRoutingTable"`
AutoOutboundsInterface *string `json:"autoOutboundsInterface"`
AutoSystemDNS bool `json:"autoSystemDNS"`
}
func (v *TunConfig) Build() (proto.Message, error) {
@@ -31,6 +32,7 @@ func (v *TunConfig) Build() (proto.Message, error) {
DNS: v.DNS,
UserLevel: v.UserLevel,
AutoSystemRoutingTable: v.AutoSystemRoutingTable,
AutoSystemDns: v.AutoSystemDNS,
}
if v.AutoOutboundsInterface != nil {
config.AutoOutboundsInterface = *v.AutoOutboundsInterface
+3
View File
@@ -312,6 +312,9 @@ func (c *VLessOutboundConfig) Build() (proto.Message, error) {
if err := json.Unmarshal(rawUser, account); err != nil {
return nil, errors.New(`VLESS users: invalid user`).Base(err)
}
// validateOutboundTransportSecurity needs to see these
c.Encryption = account.Encryption
c.Address = rec.Address
if account.Reverse != nil { // may not be reached: error json unmarshal
return nil, errors.New(`VLESS users: please use simplified outbound's config style to use "reverse"`)
}
+7 -23
View File
@@ -59,14 +59,13 @@ func (c *WireGuardPeerConfig) Build() (*wireguard.PeerConfig, error) {
type WireGuardConfig struct {
IsClient bool `json:""`
NoKernelTun bool `json:"noKernelTun"`
SecretKey string `json:"secretKey"`
Address []string `json:"address"`
Peers []*WireGuardPeerConfig `json:"peers"`
MTU int32 `json:"mtu"`
Reserved []byte `json:"reserved"`
DomainStrategy string `json:"domainStrategy"`
DNS []string `json:"remoteDNS"`
NoKernelTun bool `json:"noKernelTun"`
SecretKey string `json:"secretKey"`
Address []string `json:"address"`
Peers []*WireGuardPeerConfig `json:"peers"`
MTU int32 `json:"mtu"`
Reserved []byte `json:"reserved"`
DNS []string `json:"remoteDNS"`
}
func (c *WireGuardConfig) Build() (proto.Message, error) {
@@ -125,21 +124,6 @@ func (c *WireGuardConfig) Build() (proto.Message, error) {
}
config.Reserved = c.Reserved
switch strings.ToLower(c.DomainStrategy) {
case "forceip", "":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP
case "forceipv4":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP4
case "forceipv6":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP6
case "forceipv4v6":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP46
case "forceipv6v4":
config.DomainStrategy = wireguard.DeviceConfig_FORCE_IP64
default:
return nil, errors.New("unsupported domain strategy: ", c.DomainStrategy)
}
config.IsClient = c.IsClient
config.NoKernelTun = c.NoKernelTun
config.DNS = c.DNS
+15 -1
View File
@@ -16,6 +16,7 @@ import (
"github.com/xtls/xray-core/common/serial"
core "github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/proxy/freedom"
"github.com/xtls/xray-core/proxy/masque"
"github.com/xtls/xray-core/transport/internet"
)
@@ -32,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")
@@ -48,6 +50,7 @@ var (
"vmess": func() interface{} { return new(VMessOutboundConfig) },
"trojan": func() interface{} { return new(TrojanClientConfig) },
"hysteria": func() interface{} { return new(HysteriaClientConfig) },
"masque": func() interface{} { return new(MasqueClientConfig) },
"dns": func() interface{} { return new(DNSOutboundConfig) },
"wireguard": func() interface{} { return &WireGuardConfig{IsClient: true} },
}, "protocol", "settings")
@@ -203,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,
@@ -242,7 +248,7 @@ func validateOutboundTransportSecurity(rawConfig interface{}, senderSettings *pr
if vlessCfg.Encryption != "" && vlessCfg.Encryption != "none" {
return nil
}
if requiresTransportSecurity(vlessCfg.Vnext[0].Address) {
if requiresTransportSecurity(vlessCfg.Address) {
return errors.New("vless without TLS or other encryption is prohibited unless the server address is a private IP or domain")
}
}
@@ -338,6 +344,14 @@ func (c *OutboundDetourConfig) Build() (*core.OutboundHandlerConfig, error) {
return nil, err
}
if _, ok := ts.(*masque.ClientConfig); ok {
if ms := senderSettings.MultiplexSettings; ms != nil && ms.Enabled {
return nil, errors.New(`masque outbound does not support "mux"`)
}
} else if senderSettings.StreamSettings != nil && senderSettings.StreamSettings.ProtocolName == "masque" {
return nil, errors.New("the masque transport can only be used by the masque outbound")
}
if fc, ok := ts.(*freedom.Config); ok {
if senderSettings.StreamSettings != nil &&
senderSettings.StreamSettings.SocketSettings != nil &&
+273 -139
View File
@@ -1,15 +1,18 @@
package main
import (
"errors"
"bytes"
"flag"
"fmt"
"go/build"
"os"
"os/exec"
"path/filepath"
"runtime"
"sort"
"strings"
"sync"
"sync/atomic"
"mvdan.cc/gofumpt/format"
)
var (
@@ -23,101 +26,27 @@ var (
isFormat bool
)
// envFile returns the name of the Go environment configuration file.
// Copy from https://github.com/golang/go/blob/c4f2a9788a7be04daf931ac54382fbe2cb754938/src/cmd/go/internal/cfg/cfg.go#L150-L166
func envFile() (string, error) {
if file := os.Getenv("GOENV"); file != "" {
if file == "off" {
return "", errors.New("GOENV=off")
}
return file, nil
}
dir, err := os.UserConfigDir()
func getModuleInfo(pwd string) (modPath, langVersion string, err error) {
data, err := os.ReadFile(filepath.Join(pwd, "go.mod"))
if err != nil {
return "", err
return "", "", err
}
if dir == "" {
return "", errors.New("missing user-config dir")
}
return filepath.Join(dir, "go", "env"), nil
}
// GetRuntimeEnv returns the value of runtime environment variable,
// that is set by running following command: `go env -w key=value`.
func GetRuntimeEnv(key string) (string, error) {
file, err := envFile()
if err != nil {
return "", err
}
if file == "" {
return "", errors.New("missing runtime env file")
}
var data []byte
var runtimeEnv string
data, readErr := os.ReadFile(file)
if readErr != nil {
return "", readErr
}
envStrings := strings.Split(string(data), "\n")
for _, envItem := range envStrings {
envItem = strings.TrimSuffix(envItem, "\r")
envKeyValue := strings.Split(envItem, "=")
if len(envKeyValue) == 2 && strings.TrimSpace(envKeyValue[0]) == key {
runtimeEnv = strings.TrimSpace(envKeyValue[1])
}
}
return runtimeEnv, nil
}
// GetGOBIN returns GOBIN environment variable as a string. It will NOT be empty.
func GetGOBIN() string {
// The one set by user explicitly by `export GOBIN=/path` or `env GOBIN=/path command`
GOBIN := os.Getenv("GOBIN")
if GOBIN == "" {
var err error
// The one set by user by running `go env -w GOBIN=/path`
GOBIN, err = GetRuntimeEnv("GOBIN")
if err != nil {
// The default one that Golang uses
return filepath.Join(build.Default.GOPATH, "bin")
}
if GOBIN == "" {
return filepath.Join(build.Default.GOPATH, "bin")
}
return GOBIN
}
return GOBIN
}
func Run(binary string, args []string) ([]byte, error) {
cmd := exec.Command(binary, args...)
cmd.Env = append(cmd.Env, os.Environ()...)
output, cmdErr := cmd.CombinedOutput()
if cmdErr != nil {
return nil, cmdErr
}
return output, nil
}
func RunMany(binary string, args, files []string) bool {
fmt.Println("Processing with", binary, args, "...")
formatRequired := false
maxTasks := make(chan struct{}, runtime.NumCPU())
for _, file := range files {
maxTasks <- struct{}{}
go func(file string) {
output, err := Run(binary, append(args, file))
if err != nil {
fmt.Println(err)
} else if len(output) > 0 {
fmt.Println(string(output))
formatRequired = true
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Fields(line)
if len(fields) >= 2 {
switch fields[0] {
case "module":
modPath = fields[1]
case "go":
langVersion = "go" + strings.TrimPrefix(fields[1], "go")
}
<-maxTasks
}(file)
}
}
return formatRequired
return modPath, langVersion, nil
}
func formatGoSource(src []byte, opts format.Options) ([]byte, error) {
return format.Source(src, opts)
}
func main() {
@@ -150,26 +79,76 @@ func main() {
}
pwd := *directory
GOBIN := GetGOBIN()
binPath := os.Getenv("PATH")
pathSlice := []string{pwd, GOBIN, binPath}
binPath = strings.Join(pathSlice, string(os.PathListSeparator))
os.Setenv("PATH", binPath)
suffix := ""
if runtime.GOOS == "windows" {
suffix = ".exe"
}
gofmt := "gofumpt" + suffix
if gofmtPath, err := exec.LookPath(gofmt); err != nil {
fmt.Println("Can not find", gofmt, "in system path or current working directory.")
modPath, langVersion, modErr := getModuleInfo(pwd)
if modErr != nil {
fmt.Println("Error reading go.mod:", modErr)
os.Exit(1)
} else {
gofmt = gofmtPath
}
opts := format.Options{
LangVersion: langVersion,
ModulePath: modPath,
}
if isFormat {
fmt.Println("Formatting Go source files...")
} else if isCheck {
fmt.Println("Checking files thar are not properly formatted...")
}
jobs := make(chan string, runtime.NumCPU())
var wg sync.WaitGroup
var formatRequired atomic.Bool
var hasErrors atomic.Bool
for i := 0; i < runtime.NumCPU(); i++ {
wg.Go(func() {
for path := range jobs {
src, err := os.ReadFile(path)
if err != nil {
fmt.Fprintf(os.Stderr, "Error reading %s: %v\n", path, err)
hasErrors.Store(true)
continue
}
formatted, err := formatGoSource(src, opts)
if err != nil {
fmt.Fprintf(os.Stderr, "Error formatting %s: %v\n", path, err)
hasErrors.Store(true)
continue
}
if !bytes.Equal(src, formatted) {
var diffText []byte
if isDryrun {
newName := filepath.ToSlash(path)
oldName := newName + ".orig"
diffText = diff(oldName, src, newName, formatted)
}
if isFormat {
info, statErr := os.Stat(path)
if statErr != nil {
fmt.Fprintf(os.Stderr, "Error stating %s: %v\n", path, statErr)
hasErrors.Store(true)
continue
}
if writeErr := os.WriteFile(path, formatted, info.Mode().Perm()); writeErr != nil {
fmt.Fprintf(os.Stderr, "Error writing %s: %v\n", path, writeErr)
hasErrors.Store(true)
continue
}
}
formatRequired.Store(true)
if isDryrun && len(diffText) > 0 {
fmt.Printf("%s\n%s", path, diffText)
} else {
fmt.Println(path)
}
}
}
})
}
rawFilesSlice := make([]string, 0, 1000)
walkErr := filepath.Walk(pwd, func(path string, info os.FileInfo, err error) error {
if err != nil {
fmt.Println(err)
@@ -186,51 +165,206 @@ func main() {
!strings.HasSuffix(filename, ".pb.go") &&
!strings.Contains(dir, filepath.Join("testing", "mocks")) &&
!strings.Contains(path, filepath.Join("main", "distro", "all", "all.go")) {
rawFilesSlice = append(rawFilesSlice, path)
jobs <- path
}
return nil
})
close(jobs)
wg.Wait()
if walkErr != nil {
fmt.Println(walkErr)
os.Exit(1)
}
if isFormat {
gofmtArgs := []string{
"-l", "-e", "-w",
}
if hasErrors.Load() {
os.Exit(1)
}
fmt.Println("Formatting Go source files...")
RunMany(gofmt, gofmtArgs, rawFilesSlice)
fmt.Println("Do NOT forget to commit file changes.")
if isFormat {
if formatRequired.Load() {
fmt.Println("Do NOT forget to commit file changes.")
}
}
if isCheck {
gofmtListArgs := []string{
"-l", "-e",
}
fmt.Println("Checking files thar are not properly formatted...")
formatRequired := RunMany(gofmt, gofmtListArgs, rawFilesSlice)
if formatRequired {
if formatRequired.Load() {
fmt.Println("Format problem(s) found.")
}
if isDryrun {
if formatRequired {
gofmtShowArgs := []string{
"-d", "-e",
}
RunMany(gofmt, gofmtShowArgs, rawFilesSlice)
}
}
if formatRequired {
fmt.Println("Please run 'go install -v mvdan.cc/gofumpt@latest', then run 'go run ./infra/vformat/main.go' to format the Go source files.")
fmt.Println("Please run 'go run ./infra/vformat/main.go' to format the Go source files.")
os.Exit(1)
} else {
fmt.Println("All Go source file format check has been passed.")
}
}
}
// diff algorithm copied from mvdan.cc/gofumpt/internal/govendor/diff
type pair struct{ x, y int }
func diff(oldName string, old []byte, newName string, new []byte) []byte {
if bytes.Equal(old, new) {
return nil
}
x := diffLines(old)
y := diffLines(new)
var out bytes.Buffer
fmt.Fprintf(&out, "diff %s %s\n", oldName, newName)
fmt.Fprintf(&out, "--- %s\n", oldName)
fmt.Fprintf(&out, "+++ %s\n", newName)
var (
done pair
chunk pair
count pair
ctext []string
)
for _, m := range diffTgs(x, y) {
if m.x < done.x {
continue
}
start := m
for start.x > done.x && start.y > done.y && x[start.x-1] == y[start.y-1] {
start.x--
start.y--
}
end := m
for end.x < len(x) && end.y < len(y) && x[end.x] == y[end.y] {
end.x++
end.y++
}
for _, s := range x[done.x:start.x] {
ctext = append(ctext, "-"+s)
count.x++
}
for _, s := range y[done.y:start.y] {
ctext = append(ctext, "+"+s)
count.y++
}
const C = 3
if (end.x < len(x) || end.y < len(y)) &&
(end.x-start.x < C || (len(ctext) > 0 && end.x-start.x < 2*C)) {
for _, s := range x[start.x:end.x] {
ctext = append(ctext, " "+s)
count.x++
count.y++
}
done = end
continue
}
if len(ctext) > 0 {
n := end.x - start.x
if n > C {
n = C
}
for _, s := range x[start.x : start.x+n] {
ctext = append(ctext, " "+s)
count.x++
count.y++
}
done = pair{start.x + n, start.y + n}
if count.x > 0 {
chunk.x++
}
if count.y > 0 {
chunk.y++
}
fmt.Fprintf(&out, "@@ -%d,%d +%d,%d @@\n", chunk.x, count.x, chunk.y, count.y)
for _, s := range ctext {
out.WriteString(s)
}
count.x = 0
count.y = 0
ctext = ctext[:0]
}
if end.x >= len(x) && end.y >= len(y) {
break
}
chunk = pair{end.x - C, end.y - C}
for _, s := range x[chunk.x:end.x] {
ctext = append(ctext, " "+s)
count.x++
count.y++
}
done = end
}
return out.Bytes()
}
func diffLines(x []byte) []string {
l := strings.SplitAfter(string(x), "\n")
if l[len(l)-1] == "" {
l = l[:len(l)-1]
} else {
l[len(l)-1] += "\n\\ No newline at end of file\n"
}
return l
}
func diffTgs(x, y []string) []pair {
m := make(map[string]int)
for _, s := range x {
if c := m[s]; c > -2 {
m[s] = c - 1
}
}
for _, s := range y {
if c := m[s]; c > -8 {
m[s] = c - 4
}
}
var xi, yi, inv []int
for i, s := range y {
if m[s] == -5 {
m[s] = len(yi)
yi = append(yi, i)
}
}
for i, s := range x {
if j, ok := m[s]; ok && j >= 0 {
xi = append(xi, i)
inv = append(inv, j)
}
}
J := inv
n := len(xi)
T := make([]int, n)
L := make([]int, n)
for i := range T {
T[i] = n + 1
}
for i := 0; i < n; i++ {
k := sort.Search(n, func(k int) bool {
return T[k] >= J[i]
})
T[k] = J[i]
L[i] = k + 1
}
k := 0
for _, v := range L {
if k < v {
k = v
}
}
seq := make([]pair, 2+k)
seq[1+k] = pair{len(x), len(y)}
lastj := n
for i := n - 1; i >= 0; i-- {
if L[i] == k && J[i] < lastj {
seq[k] = pair{xi[i], yi[J[i]]}
k--
}
}
seq[0] = pair{0, 0}
return seq
}
@@ -12,6 +12,7 @@ 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/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 +89,8 @@ func extractInboundUsers(inb *core.InboundHandlerConfig) []*protocol.User {
return ty.Users
case *shadowsocks_2022.MultiUserServerConfig:
return ty.Users
case *masque.ServerConfig:
return ty.Users
default:
fmt.Println("unsupported inbound type")
}
+3
View File
@@ -41,6 +41,7 @@ import (
_ "github.com/xtls/xray-core/proxy/freedom"
_ "github.com/xtls/xray-core/proxy/http"
_ "github.com/xtls/xray-core/proxy/loopback"
_ "github.com/xtls/xray-core/proxy/masque"
_ "github.com/xtls/xray-core/proxy/shadowsocks"
_ "github.com/xtls/xray-core/proxy/socks"
_ "github.com/xtls/xray-core/proxy/trojan"
@@ -54,12 +55,14 @@ import (
_ "github.com/xtls/xray-core/transport/internet/grpc"
_ "github.com/xtls/xray-core/transport/internet/httpupgrade"
_ "github.com/xtls/xray-core/transport/internet/kcp"
_ "github.com/xtls/xray-core/transport/internet/masque"
_ "github.com/xtls/xray-core/transport/internet/reality"
_ "github.com/xtls/xray-core/transport/internet/splithttp"
_ "github.com/xtls/xray-core/transport/internet/tcp"
_ "github.com/xtls/xray-core/transport/internet/tls"
_ "github.com/xtls/xray-core/transport/internet/udp"
_ "github.com/xtls/xray-core/transport/internet/websocket"
_ "github.com/xtls/xray-core/transport/internet/xdrive"
// Transport headers
_ "github.com/xtls/xray-core/transport/internet/headers/http"
+4 -4
View File
@@ -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
+4 -8
View File
@@ -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
}
+4 -4
View File
@@ -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())
@@ -236,14 +236,14 @@ type UDPReader struct {
func (r *UDPReader) ReadFrom(p []byte) (n int, addr *net.Destination, err error) {
for {
var buf [hysteria.MaxDatagramFrameSize]byte
var packet [1500]byte
n, err := r.reader.Read(buf[:])
n, err := r.reader.Read(packet[:])
if err != nil {
return 0, nil, err
}
msg, err := ParseUDPMessage(buf[:n])
msg, err := ParseUDPMessage(packet[:n])
if err != nil {
continue
}
+2 -2
View File
@@ -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)
}
}
+1 -1
View File
@@ -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
}
+108
View File
@@ -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))
}
+49
View File
@@ -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())
}
+328
View File
@@ -0,0 +1,328 @@
package masque
import (
"context"
go_errors "errors"
"io"
"net/netip"
"slices"
"sync"
"sync/atomic"
"time"
"golang.zx2c4.com/wireguard/tun"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/common/task"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/proxy/wireguard"
"github.com/xtls/xray-core/transport"
"github.com/xtls/xray-core/transport/internet"
"github.com/xtls/xray-core/transport/internet/masque"
"github.com/xtls/xray-core/transport/internet/stat"
"github.com/xtls/xray-core/transport/internet/tls"
)
const (
establishTimeout = 10 * time.Second
retryInterval = time.Second
)
type Client struct {
server *protocol.ServerSpec
policyManager policy.Manager
remoteDNS []netip.Addr
ctx context.Context
cancel context.CancelFunc
tunnel atomic.Pointer[tunnel]
mu sync.Mutex
lastErr error
lastErrAt time.Time
}
func NewClient(ctx context.Context, config *ClientConfig) (*Client, error) {
v := core.MustFromContext(ctx)
p := v.GetFeature(policy.ManagerType()).(policy.Manager)
streamSettings := session.StreamSettingsFromContext(ctx).(*internet.MemoryStreamConfig)
if _, ok := streamSettings.ProtocolSettings.(*masque.Config); !ok {
return nil, errors.New("not masque transport")
}
if tls.ConfigFromStreamSettings(streamSettings) == nil {
return nil, errors.New(`MASQUE requires "security": "tls"`)
}
if config.Server == nil {
return nil, errors.New(`no target server found`)
}
server, err := protocol.NewServerSpecFromPB(config.Server)
if err != nil {
return nil, errors.New("failed to get server spec").Base(err)
}
dns := config.RemoteDns
if len(dns) == 0 {
dns = []string{"1.1.1.1", "1.0.0.1", "2606:4700:4700::1111", "2606:4700:4700::1001"}
}
remoteDNS := make([]netip.Addr, 0, len(dns))
for _, s := range dns {
addr, err := netip.ParseAddr(s)
if err != nil {
return nil, errors.New("invalid remote DNS server ", s).Base(err)
}
remoteDNS = append(remoteDNS, addr)
}
c := &Client{
server: server,
policyManager: p,
remoteDNS: remoteDNS,
}
c.ctx, c.cancel = context.WithCancel(context.Background())
return c, nil
}
func (c *Client) Process(ctx context.Context, link *transport.Link, dialer internet.Dialer) error {
outbounds := session.OutboundsFromContext(ctx)
ob := outbounds[len(outbounds)-1]
if !ob.Target.IsValid() {
return errors.New("target not specified")
}
ob.Name = "masque"
ob.CanSpliceCopy = 3
t, err := c.getTunnel(ctx, dialer)
if err != nil {
return errors.New("failed to establish CONNECT-IP tunnel").Base(err)
}
var newCtx context.Context
var newCancel context.CancelFunc
if session.TimeoutOnlyFromContext(ctx) {
newCtx, newCancel = context.WithCancel(context.Background())
}
sessionPolicy := c.policyManager.ForLevel(0)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, func() {
cancel()
if newCancel != nil {
newCancel()
}
}, sessionPolicy.Timeouts.ConnectionIdle)
if newCtx != nil {
ctx = newCtx
}
var reader buf.Reader
var writer buf.Writer
switch ob.Target.Network {
case net.Network_TCP:
var conn net.Conn
var err error
if sessionPolicy.Timeouts.Handshake != 0 {
timeoutCtx, timeoutCancel := context.WithTimeout(ctx, sessionPolicy.Timeouts.Handshake)
conn, err = t.tnet.DialContext(timeoutCtx, "tcp", ob.Target.NetAddr())
timeoutCancel()
} else {
conn, err = t.tnet.Dial("tcp", ob.Target.NetAddr())
}
if err != nil {
return errors.New("failed to create TCP connection").Base(err)
}
defer conn.Close()
reader = buf.NewReader(conn)
writer = buf.NewWriter(conn)
case net.Network_UDP:
conn, err := t.tnet.Dial("udp", ob.Target.NetAddr())
if err != nil {
return errors.New("failed to create UDP connection").Base(err)
}
defer conn.Close()
uc := &wireguard.UDPConnClient{
PacketConn: conn.(*net.PacketConnWrapper).PacketConn,
Dest: conn.RemoteAddr().(*net.UDPAddr),
}
reader = uc
writer = uc
default:
panic(ob.Target.Network)
}
requestFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseFunc := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDonePost := task.OnSuccess(responseFunc, task.Close(link.Writer))
if err := task.Run(ctx, requestFunc, responseDonePost); err != nil {
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
return errors.New("connection ends").Base(err)
}
return nil
}
func (c *Client) getTunnel(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
c.mu.Lock()
defer c.mu.Unlock()
if c.ctx.Err() != nil {
return nil, errors.New("closed")
}
if t := c.tunnel.Load(); t != nil {
select {
case <-t.done:
default:
return t, nil
}
}
if err := ctx.Err(); err != nil {
return nil, err
}
if c.lastErr != nil && time.Since(c.lastErrAt) < retryInterval {
return nil, c.lastErr
}
t, err := c.establish(ctx, dialer)
if err != nil {
c.lastErr, c.lastErrAt = err, time.Now()
return nil, err
}
c.lastErr = nil
c.tunnel.Store(t)
if c.ctx.Err() != nil {
if c.tunnel.CompareAndSwap(t, nil) {
t.close()
}
return nil, errors.New("closed")
}
return t, nil
}
func (c *Client) establish(ctx context.Context, dialer internet.Dialer) (*tunnel, error) {
ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), establishTimeout)
defer cancel()
defer context.AfterFunc(c.ctx, cancel)()
conn, err := dialer.Dial(ctx, c.server.Destination)
if err != nil {
return nil, err
}
mconn, ok := stat.TryUnwrapStatsConn(conn).(*masque.Conn)
if !ok {
conn.Close()
return nil, errors.New("not a CONNECT-IP connection")
}
t, err := newTunnel(conn, mconn.LocalAddrs(), c.remoteDNS)
if err != nil {
conn.Close()
return nil, err
}
errors.LogInfo(ctx, "MASQUE: tunnel established from ", mconn.LocalAddrs())
return t, nil
}
func (c *Client) Close() error {
c.cancel()
if t := c.tunnel.Swap(nil); t != nil {
t.close()
}
return nil
}
type tunnel struct {
conn stat.Connection
dev tun.Device
tnet *wireguard.Net
done chan struct{}
closeOnce sync.Once
}
func newTunnel(conn stat.Connection, local []netip.Addr, remoteDNS []netip.Addr) (*tunnel, error) {
var dns []netip.Addr
for _, addr := range remoteDNS {
if slices.ContainsFunc(local, func(l netip.Addr) bool { return l.Is4() == addr.Is4() }) {
dns = append(dns, addr)
}
}
if len(dns) == 0 {
errors.LogWarning(context.Background(), "MASQUE: no remote DNS server is reachable from the assigned addresses ", local, ", domain names will fail to resolve")
dns = remoteDNS
}
dev, tnet, _, err := wireguard.CreateNetTUN(local, dns, masque.MinPacketSize, true)
if err != nil {
return nil, err
}
t := &tunnel{
conn: conn,
dev: dev,
tnet: tnet,
done: make(chan struct{}),
}
go t.readFromTunnel()
go t.writeToTunnel()
return t, nil
}
func (t *tunnel) readFromTunnel() {
defer t.close()
b := make([]byte, buf.Size)
for {
n, err := t.conn.Read(b)
if err != nil {
if go_errors.Is(err, io.ErrShortBuffer) {
continue
}
errors.LogInfoInner(context.Background(), err, "MASQUE: tunnel closed")
return
}
t.dev.Write([][]byte{b[:n]}, 0)
}
}
func (t *tunnel) writeToTunnel() {
bufs := [][]byte{make([]byte, masque.MinPacketSize)}
sizes := []int{0}
for {
if _, err := t.dev.Read(bufs, sizes, 0); err != nil {
return
}
if _, err := t.conn.Write(bufs[0][:sizes[0]]); err != nil {
var ptb *masque.PacketTooBigError
if go_errors.As(err, &ptb) {
go t.dev.Write([][]byte{ptb.ICMP}, 0)
}
}
}
}
func (t *tunnel) close() {
t.closeOnce.Do(func() {
close(t.done)
t.conn.Close()
t.dev.Close()
})
}
func init() {
common.Must(common.RegisterConfig((*ClientConfig)(nil), func(ctx context.Context, config interface{}) (interface{}, error) {
return NewClient(ctx, config.(*ClientConfig))
}))
}
+250
View File
@@ -0,0 +1,250 @@
// Code generated by protoc-gen-go. DO NOT EDIT.
// versions:
// protoc-gen-go v1.36.11
// protoc v6.33.5
// source: proxy/masque/config.proto
package masque
import (
protocol "github.com/xtls/xray-core/common/protocol"
protoreflect "google.golang.org/protobuf/reflect/protoreflect"
protoimpl "google.golang.org/protobuf/runtime/protoimpl"
reflect "reflect"
sync "sync"
unsafe "unsafe"
)
const (
// Verify that this generated code is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(20 - protoimpl.MinVersion)
// Verify that runtime/protoimpl is sufficiently up-to-date.
_ = protoimpl.EnforceVersion(protoimpl.MaxVersion - 20)
)
type ClientConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
Server *protocol.ServerEndpoint `protobuf:"bytes,1,opt,name=server,proto3" json:"server,omitempty"`
RemoteDns []string `protobuf:"bytes,2,rep,name=remote_dns,json=remoteDns,proto3" json:"remote_dns,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *ClientConfig) Reset() {
*x = ClientConfig{}
mi := &file_proxy_masque_config_proto_msgTypes[0]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ClientConfig) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*ClientConfig) ProtoMessage() {}
func (x *ClientConfig) ProtoReflect() protoreflect.Message {
mi := &file_proxy_masque_config_proto_msgTypes[0]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use ClientConfig.ProtoReflect.Descriptor instead.
func (*ClientConfig) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{0}
}
func (x *ClientConfig) GetServer() *protocol.ServerEndpoint {
if x != nil {
return x.Server
}
return nil
}
func (x *ClientConfig) GetRemoteDns() []string {
if x != nil {
return x.RemoteDns
}
return nil
}
type Account struct {
state protoimpl.MessageState `protogen:"open.v1"`
Password string `protobuf:"bytes,1,opt,name=password,proto3" json:"password,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *Account) Reset() {
*x = Account{}
mi := &file_proxy_masque_config_proto_msgTypes[1]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *Account) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*Account) ProtoMessage() {}
func (x *Account) ProtoReflect() protoreflect.Message {
mi := &file_proxy_masque_config_proto_msgTypes[1]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use Account.ProtoReflect.Descriptor instead.
func (*Account) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{1}
}
func (x *Account) GetPassword() string {
if x != nil {
return x.Password
}
return ""
}
type ServerConfig struct {
state protoimpl.MessageState `protogen:"open.v1"`
Users []*protocol.User `protobuf:"bytes,1,rep,name=users,proto3" json:"users,omitempty"`
Address []string `protobuf:"bytes,2,rep,name=address,proto3" json:"address,omitempty"`
Mtu uint32 `protobuf:"varint,3,opt,name=mtu,proto3" json:"mtu,omitempty"`
unknownFields protoimpl.UnknownFields
sizeCache protoimpl.SizeCache
}
func (x *ServerConfig) Reset() {
*x = ServerConfig{}
mi := &file_proxy_masque_config_proto_msgTypes[2]
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
ms.StoreMessageInfo(mi)
}
func (x *ServerConfig) String() string {
return protoimpl.X.MessageStringOf(x)
}
func (*ServerConfig) ProtoMessage() {}
func (x *ServerConfig) ProtoReflect() protoreflect.Message {
mi := &file_proxy_masque_config_proto_msgTypes[2]
if x != nil {
ms := protoimpl.X.MessageStateOf(protoimpl.Pointer(x))
if ms.LoadMessageInfo() == nil {
ms.StoreMessageInfo(mi)
}
return ms
}
return mi.MessageOf(x)
}
// Deprecated: Use ServerConfig.ProtoReflect.Descriptor instead.
func (*ServerConfig) Descriptor() ([]byte, []int) {
return file_proxy_masque_config_proto_rawDescGZIP(), []int{2}
}
func (x *ServerConfig) GetUsers() []*protocol.User {
if x != nil {
return x.Users
}
return nil
}
func (x *ServerConfig) GetAddress() []string {
if x != nil {
return x.Address
}
return nil
}
func (x *ServerConfig) GetMtu() uint32 {
if x != nil {
return x.Mtu
}
return 0
}
var File_proxy_masque_config_proto protoreflect.FileDescriptor
const file_proxy_masque_config_proto_rawDesc = "" +
"\n" +
"\x19proxy/masque/config.proto\x12\x11xray.proxy.masque\x1a!common/protocol/server_spec.proto\x1a\x1acommon/protocol/user.proto\"k\n" +
"\fClientConfig\x12<\n" +
"\x06server\x18\x01 \x01(\v2$.xray.common.protocol.ServerEndpointR\x06server\x12\x1d\n" +
"\n" +
"remote_dns\x18\x02 \x03(\tR\tremoteDns\"%\n" +
"\aAccount\x12\x1a\n" +
"\bpassword\x18\x01 \x01(\tR\bpassword\"l\n" +
"\fServerConfig\x120\n" +
"\x05users\x18\x01 \x03(\v2\x1a.xray.common.protocol.UserR\x05users\x12\x18\n" +
"\aaddress\x18\x02 \x03(\tR\aaddress\x12\x10\n" +
"\x03mtu\x18\x03 \x01(\rR\x03mtuBU\n" +
"\x15com.xray.proxy.masqueP\x01Z&github.com/xtls/xray-core/proxy/masque\xaa\x02\x11Xray.Proxy.Masqueb\x06proto3"
var (
file_proxy_masque_config_proto_rawDescOnce sync.Once
file_proxy_masque_config_proto_rawDescData []byte
)
func file_proxy_masque_config_proto_rawDescGZIP() []byte {
file_proxy_masque_config_proto_rawDescOnce.Do(func() {
file_proxy_masque_config_proto_rawDescData = protoimpl.X.CompressGZIP(unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)))
})
return file_proxy_masque_config_proto_rawDescData
}
var file_proxy_masque_config_proto_msgTypes = make([]protoimpl.MessageInfo, 3)
var file_proxy_masque_config_proto_goTypes = []any{
(*ClientConfig)(nil), // 0: xray.proxy.masque.ClientConfig
(*Account)(nil), // 1: xray.proxy.masque.Account
(*ServerConfig)(nil), // 2: xray.proxy.masque.ServerConfig
(*protocol.ServerEndpoint)(nil), // 3: xray.common.protocol.ServerEndpoint
(*protocol.User)(nil), // 4: xray.common.protocol.User
}
var file_proxy_masque_config_proto_depIdxs = []int32{
3, // 0: xray.proxy.masque.ClientConfig.server:type_name -> xray.common.protocol.ServerEndpoint
4, // 1: xray.proxy.masque.ServerConfig.users:type_name -> xray.common.protocol.User
2, // [2:2] is the sub-list for method output_type
2, // [2:2] is the sub-list for method input_type
2, // [2:2] is the sub-list for extension type_name
2, // [2:2] is the sub-list for extension extendee
0, // [0:2] is the sub-list for field type_name
}
func init() { file_proxy_masque_config_proto_init() }
func file_proxy_masque_config_proto_init() {
if File_proxy_masque_config_proto != nil {
return
}
type x struct{}
out := protoimpl.TypeBuilder{
File: protoimpl.DescBuilder{
GoPackagePath: reflect.TypeOf(x{}).PkgPath(),
RawDescriptor: unsafe.Slice(unsafe.StringData(file_proxy_masque_config_proto_rawDesc), len(file_proxy_masque_config_proto_rawDesc)),
NumEnums: 0,
NumMessages: 3,
NumExtensions: 0,
NumServices: 0,
},
GoTypes: file_proxy_masque_config_proto_goTypes,
DependencyIndexes: file_proxy_masque_config_proto_depIdxs,
MessageInfos: file_proxy_masque_config_proto_msgTypes,
}.Build()
File_proxy_masque_config_proto = out.File
file_proxy_masque_config_proto_goTypes = nil
file_proxy_masque_config_proto_depIdxs = nil
}
+25
View File
@@ -0,0 +1,25 @@
syntax = "proto3";
package xray.proxy.masque;
option csharp_namespace = "Xray.Proxy.Masque";
option go_package = "github.com/xtls/xray-core/proxy/masque";
option java_package = "com.xray.proxy.masque";
option java_multiple_files = true;
import "common/protocol/server_spec.proto";
import "common/protocol/user.proto";
message ClientConfig {
xray.common.protocol.ServerEndpoint server = 1;
repeated string remote_dns = 2;
}
message Account {
string password = 1;
}
message ServerConfig {
repeated xray.common.protocol.User users = 1;
repeated string address = 2;
uint32 mtu = 3;
}
+80
View File
@@ -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)
}
+59
View File
@@ -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)
}
}
+550
View File
@@ -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))
}))
}
+328
View File
@@ -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)
}
}
+15
View File
@@ -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
+2 -2
View File
@@ -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 {
+2 -2
View File
@@ -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()
+3 -3
View File
@@ -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)}
+55
View File
@@ -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")
}
+12 -4
View File
@@ -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),
}
}
+208 -126
View File
@@ -2,17 +2,11 @@ package shadowsocks_2022
import (
"context"
"io"
"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"
@@ -20,7 +14,10 @@ import (
"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/task"
"github.com/xtls/xray-core/common/utils"
"github.com/xtls/xray-core/core"
"github.com/xtls/xray-core/features/policy"
"github.com/xtls/xray-core/features/routing"
"github.com/xtls/xray-core/transport/internet/stat"
)
@@ -32,10 +29,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 +46,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 +85,181 @@ 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()
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)
}
var salt [32]byte
saltSlice := salt[:i.method.KeySaltLength]
if _, err := io.ReadFull(conn, saltSlice); err != nil {
return err
}
if !i.saltFilter.Check(salt) {
return ErrSaltNotUnique
}
sessionKey := DeriveSessionSubKey(i.psk, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil {
return err
}
reader := NewStreamReader(conn, aead)
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination
writer, err := WriteTCPResponse(conn, i.method, i.psk, saltSlice, nil)
if err != nil {
return err
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: i.user.Email,
})
errors.LogInfo(ctx, "tunneling request to ", dest)
link, err := dispatcher.Dispatch(ctx, dest)
if err != nil {
return err
}
if len(reqHeader.EarlyData) > 0 {
earlyBuf := buf.New()
earlyBuf.Write(reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err
}
}
sessionPolicy = i.policyManager.ForLevel(uint32(i.user.Level))
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *Inbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
buf.ReleaseMulti(mb)
return err
}
for _, b := range mb {
decoded, err := i.udpCodec.DecodePacket(b.Bytes())
if err != nil {
buf.ReleaseMulti(mb)
return singbridge.ReturnError(err)
b.Release()
continue
}
for _, buffer := range mb {
packet := B.As(buffer.Bytes()).ToOwned()
buffer.Release()
err = i.service.NewPacket(ctx, pc, packet, metadata)
entry, ok := udpConns.Load(decoded.SessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: decoded.Destination,
Status: log.AccessAccepted,
Email: i.user.Email,
})
link, err := dispatcher.Dispatch(sessCtx, decoded.Destination)
if err != nil {
packet.Release()
buf.ReleaseMulti(mb)
return err
cancel()
b.Release()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(uint32(i.user.Level))
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(decoded.SessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(decoded.SessionID, newEntry)
if loaded {
// Another goroutine/packet beat us to storing, terminate our redundant link
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(sessID uint64, dest net.Destination, cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
encPacket, err := i.udpCodec.EncodeServerPacket(sessID, dest, rb.Bytes())
rb.Release()
if err != nil {
continue
}
_, _ = conn.Write(encPacket)
}
}
}(decoded.SessionID, decoded.Destination, entry)
}
}
entry.timer.Update()
payloadBuf := buf.New()
payloadBuf.Write(decoded.Payload)
b.Release()
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{payloadBuf})
}
}
}
func (i *Inbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
inbound.User = &protocol.MemoryUser{
Email: i.email,
Level: uint32(i.level),
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: i.email,
})
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
return singbridge.CopyConn(ctx, nil, link, conn)
}
func (i *Inbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
inbound.User = &protocol.MemoryUser{
Email: i.email,
Level: uint32(i.level),
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: i.email,
})
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
outConn := &singbridge.PacketConnWrapper{
Reader: link.Reader,
Writer: link.Writer,
Dest: destination,
T: signal.CancelAfterInactivity(ctx, func() {
common.Interrupt(link.Reader)
}, 300*time.Second),
}
return bufio.CopyPacketConn(ctx, conn, outConn)
}
func (i *Inbound) NewError(ctx context.Context, err error) {
if E.IsClosed(err) {
return
}
errors.LogWarning(ctx, err.Error())
}
type natPacketConn struct {
net.Conn
}
func (c *natPacketConn) ReadPacket(buffer *B.Buffer) (addr M.Socksaddr, err error) {
_, err = buffer.ReadFrom(c)
return
}
func (c *natPacketConn) WritePacket(buffer *B.Buffer, addr M.Socksaddr) error {
_, err := buffer.WriteTo(c)
return err
}
+401 -182
View File
@@ -2,21 +2,17 @@ package shadowsocks_2022
import (
"context"
"encoding/base64"
"crypto/cipher"
"encoding/binary"
"io"
"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"
@@ -24,8 +20,11 @@ import (
"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/task"
"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 +37,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 +57,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 +193,317 @@ 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. Read Request Salt (16 or 32 bytes)
var salt [32]byte
saltSlice := salt[:i.method.KeySaltLength]
if _, err := io.ReadFull(conn, saltSlice); err != nil {
return err
}
if network == net.Network_TCP {
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)
if !i.saltFilter.Check(salt) {
return ErrSaltNotUnique
}
// 2. Read Extended Identity Header (16 bytes)
var eih [AESBlockSize]byte
if _, err := io.ReadFull(conn, eih[:]); err != nil {
return err
}
// Decrypt EIH with IdentitySubKey derived from masterPSK and salt
identitySubkey := DeriveIdentitySubKey(i.masterPSK, saltSlice, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil {
return err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih[:])
// Lookup user
user, ok := i.usersByHash.Load(decryptedHash)
if !ok || user == nil {
return ErrInvalidRequest
}
userPSK := user.Account.(*MemoryAccount).Key
// 3. Derive Session Subkey using matched user's PSK
sessionKey := DeriveSessionSubKey(userPSK, saltSlice, i.method.KeySaltLength)
aead, err := i.method.NewAEAD(sessionKey)
if err != nil {
return err
}
reader := NewStreamReader(conn, aead)
// 4 & 5. Read Client Request Header
reqHeader, err := ReadClientRequestHeader(conn, reader)
if err != nil {
return err
}
conn.SetReadDeadline(time.Time{})
dest := reqHeader.Destination
// 6. Send Server Response Handshake
writer, err := WriteTCPResponse(conn, i.method, userPSK, saltSlice, nil)
if err != nil {
return err
}
// 7. Dispatch Connection to Xray routing with matched User
inbound := session.InboundFromContext(ctx)
inbound.User = user
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: user.Email,
})
errors.LogInfo(ctx, "tunneling request to ", dest, " for user ", user.Email)
link, err := dispatcher.Dispatch(ctx, dest)
if err != nil {
return err
}
if len(reqHeader.EarlyData) > 0 {
earlyBuf := buf.New()
earlyBuf.Write(reqHeader.EarlyData)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{earlyBuf}); err != nil {
return err
}
}
sessionPolicy = i.policyManager.ForLevel(user.Level)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(reader, link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *MultiUserInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
buf.ReleaseMulti(mb)
return err
}
for _, b := range mb {
// In multi-user UDP:
// Packet header is 16 bytes: Encrypted(SessionID + PacketID)
// Followed by 16 bytes EIH
packetBytes := b.Bytes()
if len(packetBytes) < 32+1+8+2 {
b.Release()
continue
}
for _, buffer := range mb {
packet := B.As(buffer.Bytes()).ToOwned()
buffer.Release()
err = i.service.NewPacket(ctx, pc, packet, metadata)
var rawHeader [16]byte
i.udpMasterCipher.Decrypt(rawHeader[:], packetBytes[:16])
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
// Replay protection & session lookup
sessionItem := i.udpSessions.GetOrCreate(sessionID)
sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
b.Release()
continue
}
var userPSK []byte
var currentUser *protocol.MemoryUser
if sessionItem.User != nil {
currentUser = sessionItem.User
userPSK = sessionItem.UserPSK
sessionItem.Unlock()
} else {
sessionItem.Unlock()
// Decrypt EIH
identitySubkey := DeriveIdentitySubKey(i.masterPSK, rawHeader[:8], i.method.KeySaltLength)
idBlock, err := i.method.NewBlock(identitySubkey)
if err != nil {
packet.Release()
buf.ReleaseMulti(mb)
return err
b.Release()
continue
}
var decryptedHash [16]byte
idBlock.Decrypt(decryptedHash[:], packetBytes[16:32])
user, ok := i.usersByHash.Load(decryptedHash)
if !ok || user == nil {
b.Release()
continue
}
currentUser = user
userPSK = user.Account.(*MemoryAccount).Key
sessionItem.Lock()
sessionItem.User = user
sessionItem.UserPSK = userPSK
sessionItem.Unlock()
}
// Decrypt Body (with AEAD caching per session)
bodyAead := sessionItem.GetRemoteCipher()
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(userPSK, rawHeader[:8], i.method.KeySaltLength)
var err error
bodyAead, err = i.method.NewAEAD(bodyKey)
if err != nil {
b.Release()
continue
}
sessionItem.SetRemoteCipher(bodyAead)
}
bodyNonce := rawHeader[4:16]
bodyCipher := packetBytes[32:]
bodyPlain, err := bodyAead.Open(nil, bodyNonce, bodyCipher, nil)
b.Release()
if err != nil || len(bodyPlain) < 1+8+2 {
continue
}
sessionItem.Lock()
sessionItem.Window.Add(packetID)
sessionItem.Unlock()
if bodyPlain[0] != HeaderTypeClient {
continue
}
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := time.Now().Unix() - int64(epoch)
if diff < -30 || diff > 30 {
continue
}
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[9:11]))
offset := 11 + paddingLen
if len(bodyPlain) < offset {
continue
}
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
if err != nil {
continue
}
payload := bodyPlain[offset+addrLen:]
entry, ok := udpConns.Load(sessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
inbound.User = currentUser
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: currentUser.Email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
cancel()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(currentUser.Level)
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(sessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
if loaded {
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(sessID uint64, uPSK []byte, d net.Destination, cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
encPacket, err := i.encodeServerUDPPacket(sessID, uPSK, d, rb.Bytes())
rb.Release()
if err != nil {
continue
}
_, _ = conn.Write(encPacket)
}
}
}(sessionID, userPSK, dest, entry)
}
}
entry.timer.Update()
pBuf := buf.New()
pBuf.Write(payload)
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{pBuf})
}
}
}
func (i *MultiUserInbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
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
func (i *MultiUserInbound) encodeServerUDPPacket(clientSessionID uint64, userPSK []byte, dest net.Destination, payload []byte) ([]byte, error) {
sessionItem := i.udpSessions.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(i.method, i.udpMasterCipher, nil, userPSK); err != nil {
return nil, 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())
return sessionItem.EncodeServerPacket(i.method, clientSessionID, dest, payload)
}
+266 -130
View File
@@ -2,18 +2,12 @@ package shadowsocks_2022
import (
"context"
"crypto/cipher"
"encoding/binary"
"io"
"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"
@@ -22,8 +16,11 @@ import (
"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/task"
"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"
)
@@ -34,10 +31,22 @@ func init() {
}))
}
type relayDest struct {
destination net.Destination
email string
level uint32
key []byte
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
rawDestinations []*RelayDestination
policyManager policy.Manager
}
func NewRelayServer(ctx context.Context, config *RelayServerConfig) (*RelayInbound, error) {
@@ -48,39 +57,63 @@ 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),
rawDestinations: config.Destinations,
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),
key: destKey,
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 +125,206 @@ 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 Salt + Outer EIH
needed := i.method.KeySaltLength + AESBlockSize
var headerBuf [48]byte
headerSlice := headerBuf[:needed]
if _, err := io.ReadFull(conn, headerSlice); err != nil {
return err
}
if network == net.Network_TCP {
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)
salt := headerSlice[:i.method.KeySaltLength]
eih := headerSlice[i.method.KeySaltLength:]
identitySubkey := DeriveIdentitySubKey(i.relayPSK, salt, i.method.KeySaltLength)
block, err := i.method.NewBlock(identitySubkey)
if err != nil {
return err
}
var decryptedHash [AESBlockSize]byte
block.Decrypt(decryptedHash[:], eih)
targetDest, ok := i.destinations[decryptedHash]
if !ok {
return ErrInvalidRequest
}
conn.SetReadDeadline(time.Time{})
inbound := session.InboundFromContext(ctx)
inbound.User = &protocol.MemoryUser{
Email: targetDest.email,
Level: targetDest.level,
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: targetDest.destination,
Status: log.AccessAccepted,
Email: targetDest.email,
})
errors.LogInfo(ctx, "relaying connection to ", targetDest.destination)
link, err := dispatcher.Dispatch(ctx, targetDest.destination)
if err != nil {
return err
}
// Unwrap outer EIH: send client salt to next hop, stripping this hop's EIH
saltBuf := buf.New()
saltBuf.Write(salt)
if err := link.Writer.WriteMultiBuffer(buf.MultiBuffer{saltBuf}); err != nil {
return err
}
sessionPolicy = i.policyManager.ForLevel(targetDest.level)
ctx, cancel := context.WithCancel(ctx)
timer := signal.CancelAfterInactivity(ctx, cancel, sessionPolicy.Timeouts.ConnectionIdle)
ctx = policy.ContextWithBufferPolicy(ctx, sessionPolicy.Buffer)
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
return buf.Copy(buf.NewReader(conn), link.Writer, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
return buf.Copy(link.Reader, buf.NewWriter(conn), buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
return task.Run(ctx, requestDone, responseDoneAndCloseWriter)
}
func (i *RelayInbound) processUDP(ctx context.Context, conn stat.Connection, dispatcher routing.Dispatcher) error {
udpConns := utils.NewTypedSyncMap[uint64, *udpConnEntry]()
defer func() {
udpConns.Range(func(key uint64, entry *udpConnEntry) bool {
entry.timer.SetTimeout(0)
return true
})
}()
reader := buf.NewReader(conn)
for {
mb, err := reader.ReadMultiBuffer()
if err != nil {
buf.ReleaseMulti(mb)
return err
}
for _, b := range mb {
data := b.Bytes()
if len(data) < 2*AESBlockSize {
b.Release()
continue
}
for _, buffer := range mb {
packet := B.As(buffer.Bytes()).ToOwned()
buffer.Release()
err = i.service.NewPacket(ctx, pc, packet, metadata)
var packetHeader [AESBlockSize]byte
i.relayBlock.Decrypt(packetHeader[:], data[:AESBlockSize])
var eiHeader [AESBlockSize]byte
i.relayBlock.Decrypt(eiHeader[:], data[AESBlockSize:2*AESBlockSize])
for idx := 0; idx < AESBlockSize; idx++ {
eiHeader[idx] ^= packetHeader[idx]
}
targetDest, ok := i.destinations[eiHeader]
if !ok {
b.Release()
continue
}
// Extract sessionID from raw packetHeader for session-level link caching before re-encrypting
sessionID := binary.BigEndian.Uint64(packetHeader[:8])
// Re-encrypt packetHeader with next hop block cipher
targetDest.blockCipher.Encrypt(packetHeader[:], packetHeader[:])
// Strip outer EIH: replace second block with re-encrypted packetHeader and advance
copy(data[AESBlockSize:2*AESBlockSize], packetHeader[:])
b.Advance(int32(AESBlockSize))
dest := targetDest.destination
dest.Network = net.Network_UDP
entry, ok := udpConns.Load(sessionID)
if !ok {
sessCtx, cancel := context.WithCancel(ctx)
inbound := session.InboundFromContext(sessCtx)
inbound.User = &protocol.MemoryUser{
Email: targetDest.email,
Level: targetDest.level,
}
sessCtx = log.ContextWithAccessMessage(sessCtx, &log.AccessMessage{
From: conn.RemoteAddr(),
To: dest,
Status: log.AccessAccepted,
Email: targetDest.email,
})
link, err := dispatcher.Dispatch(sessCtx, dest)
if err != nil {
packet.Release()
buf.ReleaseMulti(mb)
return err
cancel()
b.Release()
continue
}
newEntry := &udpConnEntry{
link: link,
cancel: cancel,
}
sessionPolicy := i.policyManager.ForLevel(targetDest.level)
newEntry.timer = signal.CancelAfterInactivity(sessCtx, func() {
udpConns.Delete(sessionID)
common.Interrupt(link.Reader)
common.Interrupt(link.Writer)
cancel()
}, sessionPolicy.Timeouts.ConnectionIdle)
actual, loaded := udpConns.LoadOrStore(sessionID, newEntry)
if loaded {
newEntry.timer.SetTimeout(0)
entry = actual
} else {
entry = newEntry
go func(cEntry *udpConnEntry) {
defer func() {
cEntry.timer.SetTimeout(0)
}()
for {
resMb, err := cEntry.link.Reader.ReadMultiBuffer()
if err != nil {
return
}
cEntry.timer.Update()
for _, rb := range resMb {
_, _ = conn.Write(rb.Bytes())
rb.Release()
}
}
}(entry)
}
}
entry.timer.Update()
_ = entry.link.Writer.WriteMultiBuffer(buf.MultiBuffer{b})
}
}
}
func (i *RelayInbound) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
userInt, _ := A.UserFromContext[int](ctx)
user := i.destinations[userInt]
inbound.User = &protocol.MemoryUser{
Email: user.Email,
Level: uint32(user.Level),
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: user.Email,
})
errors.LogInfo(ctx, "tunnelling request to tcp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_TCP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
return singbridge.CopyConn(ctx, nil, link, conn)
}
func (i *RelayInbound) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error {
inbound := session.InboundFromContext(ctx)
userInt, _ := A.UserFromContext[int](ctx)
user := i.destinations[userInt]
inbound.User = &protocol.MemoryUser{
Email: user.Email,
Level: uint32(user.Level),
}
ctx = log.ContextWithAccessMessage(ctx, &log.AccessMessage{
From: metadata.Source,
To: metadata.Destination,
Status: log.AccessAccepted,
Email: user.Email,
})
errors.LogInfo(ctx, "tunnelling request to udp:", metadata.Destination)
dispatcher := session.DispatcherFromContext(ctx)
destination, err := singbridge.ToDestination(metadata.Destination, net.Network_UDP)
if err != nil {
return err
}
link, err := dispatcher.Dispatch(ctx, destination)
if err != nil {
return err
}
outConn := &singbridge.PacketConnWrapper{
Reader: link.Reader,
Writer: link.Writer,
Dest: destination,
T: signal.CancelAfterInactivity(ctx, func() {
common.Interrupt(link.Reader)
}, 300*time.Second),
}
return bufio.CopyPacketConn(ctx, conn, outConn)
}
func (i *RelayInbound) NewError(ctx context.Context, err error) {
if E.IsClosed(err) {
return
}
errors.LogWarning(ctx, err.Error())
}
+63
View File
@@ -0,0 +1,63 @@
package shadowsocks_2022
import (
"encoding/base64"
"strings"
"lukechampine.com/blake3"
)
const (
ContextSessionSubKey = "shadowsocks 2022 session subkey"
ContextIdentitySubKey = "shadowsocks 2022 identity subkey"
)
// ParseKey decodes a base64 or raw PSK key string and validates its length
func ParseKey(key string, keyLength int) ([]byte, error) {
raw, err := base64.StdEncoding.DecodeString(key)
if err != nil {
raw = []byte(key)
}
if len(raw) != keyLength {
return nil, ErrBadKey
}
return raw, nil
}
func ParsePSKList(password string, keyLength int) ([][]byte, error) {
parts := strings.Split(password, ":")
pskList := make([][]byte, len(parts))
for i, part := range parts {
norm, err := ParseKey(part, keyLength)
if err != nil {
return nil, err
}
pskList[i] = norm
}
return pskList, nil
}
func deriveSubKey(ctx string, psk, salt []byte, keyLength int) []byte {
var keyMaterial [64]byte
kmLen := len(psk) + len(salt)
copy(keyMaterial[:], psk)
copy(keyMaterial[len(psk):], salt)
out := make([]byte, keyLength)
blake3.DeriveKey(out, ctx, keyMaterial[:kmLen])
return out
}
func DeriveSessionSubKey(psk, salt []byte, keyLength int) []byte {
return deriveSubKey(ContextSessionSubKey, psk, salt, keyLength)
}
func DeriveIdentitySubKey(psk, salt []byte, keyLength int) []byte {
return deriveSubKey(ContextIdentitySubKey, psk, salt, keyLength)
}
func DeriveUserPSKHash(userPSK []byte) [AESBlockSize]byte {
h := blake3.Sum512(userPSK)
var out [AESBlockSize]byte
copy(out[:], h[:AESBlockSize])
return out
}
+147 -90
View File
@@ -2,21 +2,20 @@ package shadowsocks_2022
import (
"context"
"crypto/rand"
"io"
"time"
shadowsocks "github.com/sagernet/sing-shadowsocks"
"github.com/sagernet/sing-shadowsocks/shadowaead_2022"
C "github.com/sagernet/sing/common"
B "github.com/sagernet/sing/common/buf"
"github.com/sagernet/sing/common/bufio"
N "github.com/sagernet/sing/common/network"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/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 +27,47 @@ 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)
}
finalPSK := pskList[len(pskList)-1]
udpCodec, err := NewUDPPacketCodec(method, finalPSK)
if err != nil {
return nil, errors.New("failed to create udp packet codec").Base(err)
}
v := core.MustFromContext(ctx)
return &Outbound{
server: net.Destination{
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 +82,123 @@ 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)
bufferedWriter := buf.NewBufferedWriter(buf.NewWriter(conn))
bodyWriter, err := WriteTCPRequest(bufferedWriter, o.method, o.pskList, destination, clientSaltSlice, nil)
if err != nil {
return errors.New("failed to write request").Base(err)
}
if err = buf.CopyOnceTimeout(link.Reader, bodyWriter, time.Millisecond*100); err != nil && err != buf.ErrNotTimeoutReader && err != buf.ErrReadTimeout {
return errors.New("failed to write A request payload").Base(err)
}
if err := bufferedWriter.SetBuffered(false); err != nil {
return err
}
return buf.Copy(link.Reader, bodyWriter, buf.UpdateActivity(timer))
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
responseReader, err := ReadTCPResponse(conn, o.method, o.finalPSK, clientSaltSlice)
if err != nil {
return err
}
return buf.Copy(responseReader, link.Writer, buf.UpdateActivity(timer))
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
if err := task.Run(ctx, requestDone, responseDoneAndCloseWriter); err != nil {
return errors.New("connection ends").Base(err)
}
return nil
}
if network == net.Network_UDP {
requestDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.DownlinkOnly)
writer := &UDPWriter{
Writer: conn,
Destination: destination,
Codec: o.udpCodec,
}
if err := buf.Copy(link.Reader, writer, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to transport all UDP request").Base(err)
}
return nil
}
responseDone := func() error {
defer timer.SetTimeout(sessionPolicy.Timeouts.UplinkOnly)
reader := &UDPReader{
Reader: conn,
Codec: o.udpCodec,
}
if err := buf.Copy(reader, link.Writer, buf.UpdateActivity(timer)); err != nil {
return errors.New("failed to transport all UDP response").Base(err)
}
return nil
}
responseDoneAndCloseWriter := task.OnSuccess(responseDone, task.Close(link.Writer))
if err := task.Run(ctx, requestDone, responseDoneAndCloseWriter); err != nil {
return errors.New("connection ends").Base(err)
}
return nil
}
return errors.New("unsupported network: ", network)
}
+512
View File
@@ -0,0 +1,512 @@
package shadowsocks_2022
import (
"crypto/cipher"
"crypto/rand"
"encoding/binary"
"io"
"math"
mrand "math/rand/v2"
"sync/atomic"
"time"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
)
type UDPCodec struct {
method *CipherMethod
psk []byte
blockCipher cipher.Block
chachaCipher cipher.AEAD
clientBodyCipher cipher.AEAD
clientSessionID uint64
nextPacketID atomic.Uint64
sessions *UDPSessionManager
}
type (
UDPPacketCodec = UDPCodec
UDPServerCodec = UDPCodec
)
func newUDPCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
c := &UDPCodec{
method: method,
psk: psk,
}
var err error
if method.IsChaCha {
c.chachaCipher, err = method.NewUDPCipher(psk)
} else {
c.blockCipher, err = method.NewBlock(psk)
}
if err != nil {
return nil, err
}
return c, nil
}
func NewUDPPacketCodec(method *CipherMethod, psk []byte) (*UDPCodec, error) {
c, err := newUDPCodec(method, psk)
if err != nil {
return nil, err
}
var sessID [8]byte
if _, err := io.ReadFull(rand.Reader, sessID[:]); err != nil {
return nil, err
}
c.clientSessionID = binary.BigEndian.Uint64(sessID[:])
if !method.IsChaCha {
clientBodyKey := DeriveSessionSubKey(psk, sessID[:], method.KeySaltLength)
c.clientBodyCipher, err = method.NewAEAD(clientBodyKey)
if err != nil {
return nil, err
}
}
return c, nil
}
func NewUDPServerCodec(method *CipherMethod, psk []byte, sessionTimeout time.Duration) (*UDPCodec, error) {
c, err := newUDPCodec(method, psk)
if err != nil {
return nil, err
}
c.sessions = NewUDPSessionManager(sessionTimeout)
return c, nil
}
func (c *UDPCodec) EncodeClientPacket(dest net.Destination, payload []byte) (*buf.Buffer, error) {
packetID := c.nextPacketID.Add(1)
sessID := c.clientSessionID
// Padding determination (e.g. DNS port 53 disguise)
var paddingLen int
if dest.Port == 53 && len(payload) < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength-len(payload)) + 1
}
addrPortLen := AddrPortLength(dest)
if c.method.IsChaCha {
// ChaCha20 mode: 24-byte nonce + plaintext header (27B) + padding + dest + payload + AEAD tag (16B)
totalLen := PacketNonceSize + 27 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var nonce [PacketNonceSize]byte
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(nonce[:])
var hdr [16 + 1 + 8 + 2]byte
binary.BigEndian.PutUint64(hdr[0:8], sessID)
binary.BigEndian.PutUint64(hdr[8:16], packetID)
hdr[16] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[25:27], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[PacketNonceSize:]
outBuf.Extend(int32(c.chachaCipher.Overhead()))
c.chachaCipher.Seal(plainBytes[:0], nonce[:], plainBytes, nil)
return outBuf, nil
}
// AES mode:
// 16B Encrypted Header + (11B header + padding + dest + payload + 16B AEAD tag)
totalLen := 16 + 11 + paddingLen + addrPortLen + len(payload) + AEADTagSize
if totalLen > buf.Size {
return nil, ErrPacketTooLarge
}
outBuf := buf.New()
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], sessID)
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
var encryptedHeader [16]byte
c.blockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
outBuf.Write(encryptedHeader[:])
bodyAead := c.clientBodyCipher
var hdr [1 + 8 + 2]byte
hdr[0] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[9:11], uint16(paddingLen))
outBuf.Write(hdr[:])
if paddingLen > 0 {
outBuf.Write(zeroPadding[:paddingLen])
}
if err := WriteAddressPort(outBuf, dest); err != nil {
outBuf.Release()
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()[16:]
bodyNonce := rawHeader[4:16]
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
return outBuf, nil
}
type DecodedUDPPacket struct {
SessionID uint64
PacketID uint64
HeaderType byte
Timestamp uint64
Destination net.Destination
Payload []byte
}
func parseAddressPort(data []byte) (net.Destination, int, error) {
if len(data) < 1 {
return net.Destination{}, 0, ErrPacketTooShort
}
switch data[0] {
case 1: // IPv4
if len(data) < 1+4+2 {
return net.Destination{}, 0, ErrPacketTooShort
}
ip := net.IPAddress(data[1:5])
port := binary.BigEndian.Uint16(data[5:7])
return net.UDPDestination(ip, net.Port(port)), 7, nil
case 4: // IPv6
if len(data) < 1+16+2 {
return net.Destination{}, 0, ErrPacketTooShort
}
ip := net.IPAddress(data[1:17])
port := binary.BigEndian.Uint16(data[17:19])
return net.UDPDestination(ip, net.Port(port)), 19, nil
case 3: // Domain
if len(data) < 2 {
return net.Destination{}, 0, ErrPacketTooShort
}
domainLen := int(data[1])
if len(data) < 2+domainLen+2 {
return net.Destination{}, 0, ErrPacketTooShort
}
domain := string(data[2 : 2+domainLen])
port := binary.BigEndian.Uint16(data[2+domainLen : 2+domainLen+2])
return net.UDPDestination(net.DomainAddress(domain), net.Port(port)), 2 + domainLen + 2, nil
default:
return net.Destination{}, 0, errors.New("unknown address type")
}
}
func parsePlainUDPPacket(sessionID, packetID uint64, bodyPlain []byte) (DecodedUDPPacket, error) {
if len(bodyPlain) < 1+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
headerType := bodyPlain[0]
epoch := binary.BigEndian.Uint64(bodyPlain[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
if diff > 30 {
return DecodedUDPPacket{}, ErrBadTimestamp
}
offset := 9
if headerType == HeaderTypeServer {
if len(bodyPlain) < offset+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
offset += 8 // skip clientSessionID
}
paddingLen := int(binary.BigEndian.Uint16(bodyPlain[offset : offset+2]))
offset += 2
if len(bodyPlain) < offset+paddingLen {
return DecodedUDPPacket{}, ErrNoPadding
}
offset += paddingLen
dest, addrLen, err := parseAddressPort(bodyPlain[offset:])
if err != nil {
return DecodedUDPPacket{}, err
}
payload := bodyPlain[offset+addrLen:]
return DecodedUDPPacket{
SessionID: sessionID,
PacketID: packetID,
HeaderType: headerType,
Timestamp: epoch,
Destination: dest,
Payload: payload,
}, nil
}
func (c *UDPCodec) DecodePacket(data []byte) (DecodedUDPPacket, error) {
if len(data) < PacketMinimalHeaderSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
if c.method.IsChaCha {
if len(data) < PacketNonceSize+AEADTagSize {
return DecodedUDPPacket{}, ErrPacketTooShort
}
nonce := data[:PacketNonceSize]
ciphertext := data[PacketNonceSize:]
plain, err := c.chachaCipher.Open(ciphertext[:0], nonce, ciphertext, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt chacha udp packet").Base(err)
}
if len(plain) < 16+1+8+2 {
return DecodedUDPPacket{}, ErrPacketTooShort
}
sessionID := binary.BigEndian.Uint64(plain[:8])
packetID := binary.BigEndian.Uint64(plain[8:16])
if c.sessions != nil {
sessionItem := c.sessions.GetOrCreate(sessionID)
sessionItem.Lock()
if !sessionItem.Window.CheckAndAdd(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
}
return parsePlainUDPPacket(sessionID, packetID, plain[16:])
}
// AES mode
var rawHeader [16]byte
c.blockCipher.Decrypt(rawHeader[:], data[:16])
sessionID := binary.BigEndian.Uint64(rawHeader[:8])
packetID := binary.BigEndian.Uint64(rawHeader[8:16])
var bodyAead cipher.AEAD
var sessionItem *ServerUDPSession
if c.sessions != nil {
sessionItem = c.sessions.GetOrCreate(sessionID)
sessionItem.Lock()
if !sessionItem.Window.Check(packetID) {
sessionItem.Unlock()
return DecodedUDPPacket{}, ErrPacketIdNotUnique
}
sessionItem.Unlock()
bodyAead = sessionItem.GetRemoteCipher()
if bodyAead == nil {
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
var err error
bodyAead, err = c.method.NewAEAD(bodyKey)
if err != nil {
return DecodedUDPPacket{}, err
}
sessionItem.SetRemoteCipher(bodyAead)
}
} else {
bodyKey := DeriveSessionSubKey(c.psk, rawHeader[:8], c.method.KeySaltLength)
var err error
bodyAead, err = c.method.NewAEAD(bodyKey)
if err != nil {
return DecodedUDPPacket{}, err
}
}
bodyNonce := rawHeader[4:16]
bodyCipher := data[16:]
bodyPlain, err := bodyAead.Open(bodyCipher[:0], bodyNonce, bodyCipher, nil)
if err != nil {
return DecodedUDPPacket{}, errors.New("failed to decrypt aes udp body").Base(err)
}
if sessionItem != nil {
sessionItem.Lock()
sessionItem.Window.Add(packetID)
sessionItem.Unlock()
}
return parsePlainUDPPacket(sessionID, packetID, bodyPlain)
}
func (s *ServerUDPSession) EnsureServerState(method *CipherMethod, headerBlock cipher.Block, chachaCipher cipher.AEAD, psk []byte) error {
s.Lock()
defer s.Unlock()
if s.ServerSessionID != 0 {
return nil
}
var sidBuf [8]byte
for {
if _, err := io.ReadFull(rand.Reader, sidBuf[:]); err != nil {
return err
}
s.ServerSessionID = binary.BigEndian.Uint64(sidBuf[:])
if s.ServerSessionID != 0 {
break
}
}
if method.IsChaCha {
s.ServerChaCha = chachaCipher
} else {
s.ServerBlockCipher = headerBlock
bodyKey := DeriveSessionSubKey(psk, sidBuf[:], method.KeySaltLength)
bodyAead, err := method.NewAEAD(bodyKey)
if err != nil {
s.ServerSessionID = 0
return err
}
s.ServerCipher = bodyAead
}
return nil
}
func (s *ServerUDPSession) EncodeServerPacket(method *CipherMethod, clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
serverSessionID := s.ServerSessionID
serverPacketID := s.ServerPacketID.Add(1)
if method.IsChaCha {
var nonce [PacketNonceSize]byte
if _, err := io.ReadFull(rand.Reader, nonce[:]); err != nil {
return nil, err
}
plainBuf := buf.New()
defer plainBuf.Release()
var hdr [16 + 1 + 8 + 8 + 2]byte
binary.BigEndian.PutUint64(hdr[0:8], serverSessionID)
binary.BigEndian.PutUint64(hdr[8:16], serverPacketID)
hdr[16] = HeaderTypeServer
binary.BigEndian.PutUint64(hdr[17:25], uint64(time.Now().Unix()))
binary.BigEndian.PutUint64(hdr[25:33], clientSessionID)
binary.BigEndian.PutUint16(hdr[33:35], 0)
plainBuf.Write(hdr[:])
if err := WriteAddressPort(plainBuf, dest); err != nil {
return nil, err
}
plainBuf.Write(payload)
sealed := s.ServerChaCha.Seal(nil, nonce[:], plainBuf.Bytes(), nil)
res := make([]byte, PacketNonceSize+len(sealed))
copy(res[:PacketNonceSize], nonce[:])
copy(res[PacketNonceSize:], sealed)
return res, nil
}
// AES mode
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], serverSessionID)
binary.BigEndian.PutUint64(rawHeader[8:16], serverPacketID)
var encryptedHeader [16]byte
s.ServerBlockCipher.Encrypt(encryptedHeader[:], rawHeader[:])
bodyBuf := buf.New()
defer bodyBuf.Release()
var hdr [1 + 8 + 8 + 2]byte
hdr[0] = HeaderTypeServer
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint64(hdr[9:17], clientSessionID)
binary.BigEndian.PutUint16(hdr[17:19], 0)
bodyBuf.Write(hdr[:])
if err := WriteAddressPort(bodyBuf, dest); err != nil {
return nil, err
}
bodyBuf.Write(payload)
bodyNonce := rawHeader[4:16]
sealedBody := s.ServerCipher.Seal(nil, bodyNonce, bodyBuf.Bytes(), nil)
res := make([]byte, 16+len(sealedBody))
copy(res[:16], encryptedHeader[:])
copy(res[16:], sealedBody)
return res, nil
}
func (c *UDPCodec) EncodeServerPacket(clientSessionID uint64, dest net.Destination, payload []byte) ([]byte, error) {
sessionItem := c.sessions.GetOrCreate(clientSessionID)
if err := sessionItem.EnsureServerState(c.method, c.blockCipher, c.chachaCipher, c.psk); err != nil {
return nil, err
}
return sessionItem.EncodeServerPacket(c.method, clientSessionID, dest, payload)
}
type UDPWriter struct {
Writer io.Writer
Destination net.Destination
Codec *UDPPacketCodec
}
func (w *UDPWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
for {
mb2, b := buf.SplitFirst(mb)
mb = mb2
if b == nil {
break
}
dest := w.Destination
if b.UDP != nil {
dest = *b.UDP
}
pktBuf, err := w.Codec.EncodeClientPacket(dest, b.Bytes())
b.Release()
if err != nil {
buf.ReleaseMulti(mb)
return err
}
_, writeErr := w.Writer.Write(pktBuf.Bytes())
pktBuf.Release()
if writeErr != nil {
buf.ReleaseMulti(mb)
return writeErr
}
}
return nil
}
type UDPReader struct {
Reader io.Reader
Codec *UDPPacketCodec
}
func (r *UDPReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
for {
buffer := buf.New()
_, err := buffer.ReadFrom(r.Reader)
if err != nil {
buffer.Release()
return nil, err
}
decoded, err := r.Codec.DecodePacket(buffer.Bytes())
if err != nil {
buffer.Release()
continue
}
buffer.Clear()
buffer.Write(decoded.Payload)
dest := decoded.Destination
buffer.UDP = &dest
return buf.MultiBuffer{buffer}, nil
}
}
+271
View File
@@ -0,0 +1,271 @@
package shadowsocks_2022_test
import (
"context"
"encoding/base64"
"encoding/binary"
"errors"
gonet "net"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/features/routing"
. "github.com/xtls/xray-core/proxy/shadowsocks_2022"
"github.com/xtls/xray-core/transport"
"lukechampine.com/blake3"
)
// encodeRelayClientUDPPacket encodes a Shadowsocks-2022 UDP packet with 1 layer of EIH (Relay)
func encodeRelayClientUDPPacket(relayKey, destKey []byte, sessionID, packetID uint64, dest net.Destination, payload []byte) ([]byte, error) {
method, err := GetCipherMethod(MethodAES128GCM)
if err != nil {
return nil, err
}
relayBlock, err := method.NewBlock(relayKey)
if err != nil {
return nil, err
}
// 1. Plain packet header: sessionID (8B) + packetID (8B)
var rawHeader [16]byte
binary.BigEndian.PutUint64(rawHeader[:8], sessionID)
binary.BigEndian.PutUint64(rawHeader[8:16], packetID)
// Encrypt packetHeader under relayKey
var encPacketHeader [16]byte
relayBlock.Encrypt(encPacketHeader[:], rawHeader[:])
// 2. EI Header: blake3(destKey)[:16] ^ rawHeader
var destHash [16]byte
hash512 := blake3.Sum512(destKey)
copy(destHash[:], hash512[:16])
var eiHeader [16]byte
for i := 0; i < 16; i++ {
eiHeader[i] = destHash[i] ^ rawHeader[i]
}
var encEIHeader [16]byte
relayBlock.Encrypt(encEIHeader[:], eiHeader[:])
// 3. Payload under destination server's AEAD
bodyKey := DeriveSessionSubKey(destKey, rawHeader[:8], 16)
bodyAead, err := method.NewAEAD(bodyKey)
if err != nil {
return nil, err
}
bodyNonce := rawHeader[4:16]
outBuf := buf.New()
defer outBuf.Release()
// VarHeader: client type (1) + timestamp (8) + paddingLen (2) + padding + dest + payload
var hdr [1 + 8 + 2]byte
hdr[0] = HeaderTypeClient
binary.BigEndian.PutUint64(hdr[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(hdr[9:11], 0)
outBuf.Write(hdr[:])
if err := WriteAddressPort(outBuf, dest); err != nil {
return nil, err
}
outBuf.Write(payload)
plainBytes := outBuf.Bytes()
outBuf.Extend(int32(bodyAead.Overhead()))
bodyAead.Seal(plainBytes[:0], bodyNonce, plainBytes, nil)
// Full packet: encPacketHeader (16B) + encEIHeader (16B) + sealedBody
packet := make([]byte, 0, 32+outBuf.Len())
packet = append(packet, encPacketHeader[:]...)
packet = append(packet, encEIHeader[:]...)
packet = append(packet, outBuf.Bytes()...)
return packet, nil
}
func TestRelayUDPSessionStabilityAndDispatch(t *testing.T) {
relayKey := []byte("0123456789abcdef")
destKey := []byte("fedcba9876543210")
relayKeyB64 := base64.StdEncoding.EncodeToString(relayKey)
destKeyB64 := base64.StdEncoding.EncodeToString(destKey)
config := &RelayServerConfig{
Method: MethodAES128GCM,
Key: relayKeyB64,
Destinations: []*RelayDestination{
{
Key: destKeyB64,
Address: &net.IPOrDomain{Address: &net.IPOrDomain_Ip{Ip: []byte{127, 0, 0, 1}}},
Port: 8388,
Email: "dest@example.com",
},
},
}
inbound, err := NewRelayServer(newTestContext(), config)
if err != nil {
t.Fatalf("failed to create RelayServer: %v", err)
}
sessionID := uint64(0x1122334455667788)
dest := net.UDPDestination(net.LocalHostIP, 8388)
pkt1, err := encodeRelayClientUDPPacket(relayKey, destKey, sessionID, 1, dest, []byte("xray packet 1"))
if err != nil {
t.Fatalf("failed to encode pkt1: %v", err)
}
pkt2, err := encodeRelayClientUDPPacket(relayKey, destKey, sessionID, 2, dest, []byte("xray packet 2"))
if err != nil {
t.Fatalf("failed to encode pkt2: %v", err)
}
var dispatchCount atomic.Int32
var receivedPackets [][]byte
var mu sync.Mutex
disp := &dummyDispatcher{
onDispatch: func(ctx context.Context, d net.Destination) (*transport.Link, error) {
dispatchCount.Add(1)
linkR, linkW := gonet.Pipe()
t.Cleanup(func() {
linkW.Close()
linkR.Close()
})
link := &transport.Link{
Reader: buf.NewReader(linkR),
Writer: &customWriter{
write: func(mb buf.MultiBuffer) error {
mu.Lock()
defer mu.Unlock()
for _, b := range mb {
cpy := make([]byte, b.Len())
copy(cpy, b.Bytes())
receivedPackets = append(receivedPackets, cpy)
b.Release()
}
return nil
},
},
}
return link, nil
},
}
clientConn, serverConn := gonet.Pipe()
defer clientConn.Close()
defer serverConn.Close()
inboundConn := &dummyStatConn{Conn: serverConn}
ctx, cancel := context.WithCancel(newTestContext())
defer cancel()
go func() {
_ = inbound.Process(ctx, net.Network_UDP, inboundConn, disp)
}()
// Send Packet 1
_, err = clientConn.Write(pkt1)
if err != nil {
t.Fatalf("write pkt1 failed: %v", err)
}
time.Sleep(50 * time.Millisecond)
// Send Packet 2 (same sessionID, packetID=2)
_, err = clientConn.Write(pkt2)
if err != nil {
t.Fatalf("write pkt2 failed: %v", err)
}
time.Sleep(50 * time.Millisecond)
// Check dispatch count: For the SAME UDP session, Dispatch MUST be called exactly ONCE!
if count := dispatchCount.Load(); count != 1 {
t.Fatalf("CRITICAL BUG CONFIRMED: expected dispatchCount = 1 for same session, got %d (sessionID was corrupted by Encrypt!)", count)
}
// Verify downstream destination can decode both packets
method, err := GetCipherMethod(MethodAES128GCM)
common.Must(err)
destCodec, err := NewUDPServerCodec(method, destKey, 300*time.Second)
common.Must(err)
mu.Lock()
pkts := receivedPackets
mu.Unlock()
if len(pkts) != 2 {
t.Fatalf("expected 2 received packets at destination, got %d", len(pkts))
}
dec1, err := destCodec.DecodePacket(pkts[0])
if err != nil {
t.Fatalf("dest failed to decode packet 1: %v", err)
}
if dec1.SessionID != sessionID || dec1.PacketID != 1 || string(dec1.Payload) != "xray packet 1" {
t.Fatalf("dec1 mismatch: sess=%x, pktID=%d, payload=%s", dec1.SessionID, dec1.PacketID, string(dec1.Payload))
}
dec2, err := destCodec.DecodePacket(pkts[1])
if err != nil {
t.Fatalf("dest failed to decode packet 2: %v", err)
}
if dec2.SessionID != sessionID || dec2.PacketID != 2 || string(dec2.Payload) != "xray packet 2" {
t.Fatalf("dec2 mismatch: sess=%x, pktID=%d, payload=%s", dec2.SessionID, dec2.PacketID, string(dec2.Payload))
}
}
type customWriter struct {
write func(mb buf.MultiBuffer) error
}
func (w *customWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
return w.write(mb)
}
func (w *customWriter) Close() error {
return nil
}
func (w *customWriter) Interrupt() {}
type dummyDispatcher struct {
onDispatch func(ctx context.Context, dest net.Destination) (*transport.Link, error)
}
func (d *dummyDispatcher) Dispatch(ctx context.Context, dest net.Destination) (*transport.Link, error) {
if d.onDispatch != nil {
return d.onDispatch(ctx, dest)
}
return nil, errors.New("not handled")
}
func (d *dummyDispatcher) DispatchLink(ctx context.Context, dest net.Destination, link *transport.Link) error {
return nil
}
func (d *dummyDispatcher) Start() error { return nil }
func (d *dummyDispatcher) Close() error { return nil }
func (d *dummyDispatcher) Type() interface{} { return routing.DispatcherType() }
type dummyStatConn struct {
gonet.Conn
}
func (c *dummyStatConn) ReadMultiBuffer() (buf.MultiBuffer, error) {
b := buf.New()
_, err := b.ReadFrom(c.Conn)
return buf.MultiBuffer{b}, err
}
func (c *dummyStatConn) WriteMultiBuffer(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
if _, err := c.Conn.Write(b.Bytes()); err != nil {
return err
}
}
return nil
}
+158
View File
@@ -0,0 +1,158 @@
package shadowsocks_2022
import (
"crypto/cipher"
"sync"
"sync/atomic"
"time"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/utils"
)
const (
swBlockBitLog = 6 // 1<<6 == 64 bits
swBlockBits = 1 << swBlockBitLog // 64
swRingBlocks = 1 << 7 // 128
swBlockMask = swRingBlocks - 1 // 127
swBitMask = swBlockBits - 1 // 63
swSize = (swRingBlocks - 1) * swBlockBits // 8128
)
type SlidingWindow struct {
last uint64
ring [swRingBlocks]uint64
}
func (f *SlidingWindow) Reset() {
*f = SlidingWindow{}
}
func (f *SlidingWindow) Check(counter uint64) bool {
switch {
case counter > f.last:
return true
case f.last-counter > swSize:
return false
}
blockIndex := (counter >> swBlockBitLog) & swBlockMask
bitIndex := counter & swBitMask
return (f.ring[blockIndex]>>bitIndex)&1 == 0
}
func (f *SlidingWindow) Add(counter uint64) {
blockIndex := counter >> swBlockBitLog
if counter > f.last {
lastBlockIndex := f.last >> swBlockBitLog
diff := int(blockIndex - lastBlockIndex)
if diff > swRingBlocks {
diff = swRingBlocks
}
for i := 0; i < diff; i++ {
lastBlockIndex = (lastBlockIndex + 1) & swBlockMask
f.ring[lastBlockIndex] = 0
}
f.last = counter
}
blockIndex &= swBlockMask
bitIndex := counter & swBitMask
f.ring[blockIndex] |= 1 << bitIndex
}
func (f *SlidingWindow) CheckAndAdd(counter uint64) bool {
if !f.Check(counter) {
return false
}
f.Add(counter)
return true
}
type ServerUDPSession struct {
sync.Mutex
SessionID uint64
RemoteCipher atomic.Pointer[cipher.AEAD]
Window SlidingWindow
User *protocol.MemoryUser
UserPSK []byte
LastActive atomic.Int64 // Unix timestamp in seconds
ServerSessionID uint64
ServerPacketID atomic.Uint64
ServerCipher cipher.AEAD
ServerBlockCipher cipher.Block
ServerChaCha cipher.AEAD
}
func (s *ServerUDPSession) GetRemoteCipher() cipher.AEAD {
ptr := s.RemoteCipher.Load()
if ptr == nil {
return nil
}
return *ptr
}
func (s *ServerUDPSession) SetRemoteCipher(c cipher.AEAD) {
s.RemoteCipher.Store(&c)
}
type UDPSessionManager struct {
sessions *utils.TypedSyncMap[uint64, *ServerUDPSession]
timeout time.Duration
lastClean atomic.Int64 // Unix timestamp in seconds
}
func NewUDPSessionManager(timeout time.Duration) *UDPSessionManager {
return &UDPSessionManager{
sessions: utils.NewTypedSyncMap[uint64, *ServerUDPSession](),
timeout: timeout,
}
}
func (m *UDPSessionManager) GetOrCreate(sessionID uint64) *ServerUDPSession {
now := time.Now().Unix()
if s, ok := m.sessions.Load(sessionID); ok {
s.LastActive.Store(now)
return s
}
s := &ServerUDPSession{
SessionID: sessionID,
}
s.LastActive.Store(now)
actual, loaded := m.sessions.LoadOrStore(sessionID, s)
if loaded {
actual.LastActive.Store(now)
return actual
}
// Trigger cleanup if at least 30 seconds have passed since last cleanup
last := m.lastClean.Load()
if now-last > 30 && m.lastClean.CompareAndSwap(last, now) {
go m.cleanup(now)
}
return s
}
func (m *UDPSessionManager) cleanup(now int64) {
timeoutSec := int64(m.timeout.Seconds())
if timeoutSec <= 0 {
timeoutSec = 60
}
m.sessions.Range(func(k uint64, v *ServerUDPSession) bool {
if now-v.LastActive.Load() > timeoutSec {
m.sessions.Delete(k)
}
return true
})
}
func (m *UDPSessionManager) Delete(sessionID uint64) {
m.sessions.Delete(sessionID)
}
@@ -1 +1,50 @@
package shadowsocks_2022
import (
"context"
"sync"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/signal"
"github.com/xtls/xray-core/transport"
)
type udpConnEntry struct {
sync.Mutex
link *transport.Link
timer *signal.ActivityTimer
cancel context.CancelFunc
}
const (
HeaderTypeClient = 0
HeaderTypeServer = 1
MaxPaddingLength = 900
PacketNonceSize = 24
MaxPacketSize = 65535
RequestHeaderFixedChunkLength = 1 + 8 + 2 // Type (1B) + Timestamp (8B) + VarHeaderLen (2B)
PacketMinimalHeaderSize = 30
StreamNonceSize = 12
AESBlockSize = 16
AEADTagSize = 16
)
var zeroPadding [MaxPaddingLength]byte
const (
MethodAES128GCM = "2022-blake3-aes-128-gcm"
MethodAES256GCM = "2022-blake3-aes-256-gcm"
MethodChaCha20Poly1305 = "2022-blake3-chacha20-poly1305"
)
var (
ErrBadKey = errors.New("bad key")
ErrBadHeaderType = errors.New("bad header type")
ErrBadTimestamp = errors.New("bad timestamp")
ErrSaltNotUnique = errors.New("salt not unique")
ErrPacketIdNotUnique = errors.New("packet id not unique")
ErrPacketTooShort = errors.New("packet too short")
ErrPacketTooLarge = errors.New("packet too large")
ErrNoPadding = errors.New("bad request: missing payload or padding")
ErrInvalidRequest = errors.New("invalid request")
)
@@ -0,0 +1,362 @@
package shadowsocks_2022_test
import (
"bytes"
"context"
"crypto/rand"
"encoding/base64"
"encoding/binary"
"io"
gonet "net"
"sync"
"testing"
"time"
"github.com/google/go-cmp/cmp"
"github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/serial"
"github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/core"
. "github.com/xtls/xray-core/proxy/shadowsocks_2022"
)
func newTestContext() context.Context {
v, err := core.New(&core.Config{})
common.Must(err)
ctx := context.WithValue(context.Background(), core.XrayKey(1), v)
ctx = session.ContextWithInbound(ctx, &session.Inbound{})
return ctx
}
func generateRandomKey(size int) string {
b := make([]byte, size)
_, _ = rand.Read(b)
return base64.StdEncoding.EncodeToString(b)
}
func TestKDF(t *testing.T) {
// Test ParseKey
if _, err := ParseKey("", 16); err != ErrBadKey {
t.Fatalf("expected ErrBadKey for empty key, got %v", err)
}
shortKey := base64.StdEncoding.EncodeToString([]byte("short"))
if _, err := ParseKey(shortKey, 16); err != ErrBadKey {
t.Fatalf("expected ErrBadKey for short key, got %v", err)
}
exactKey := []byte("0123456789abcdef")
exactKeyB64 := base64.StdEncoding.EncodeToString(exactKey)
normExact, err := ParseKey(exactKeyB64, 16)
if err != nil || !bytes.Equal(normExact, exactKey) {
t.Fatalf("unexpected parsed exact key: %v, err: %v", normExact, err)
}
longKey := base64.StdEncoding.EncodeToString([]byte("0123456789abcdef_longer_key_for_testing"))
if _, err := ParseKey(longKey, 16); err != ErrBadKey {
t.Fatalf("expected ErrBadKey for long key, got %v", err)
}
// Test Session Subkey determinism
salt := []byte("random_salt_1234")
k1 := DeriveSessionSubKey(normExact, salt, 16)
k2 := DeriveSessionSubKey(normExact, salt, 16)
if !bytes.Equal(k1, k2) {
t.Fatal("DeriveSessionSubKey should be deterministic")
}
// Identity subkey must differ from session subkey with same inputs
idKey := DeriveIdentitySubKey(normExact, salt, 16)
if bytes.Equal(k1, idKey) {
t.Fatal("DeriveIdentitySubKey must differ from DeriveSessionSubKey")
}
// User PSK hash
h1 := DeriveUserPSKHash(normExact)
h2 := DeriveUserPSKHash(normExact)
if h1 != h2 {
t.Fatal("DeriveUserPSKHash should be deterministic")
}
}
func TestSlidingWindow(t *testing.T) {
var window SlidingWindow
if !window.Check(1) {
t.Fatal("packet 1 should be accepted")
}
window.Add(1)
if window.Check(1) {
t.Fatal("duplicate packet 1 should be rejected")
}
if !window.Check(100) {
t.Fatal("packet 100 should be accepted")
}
window.Add(100)
if window.Check(100) {
t.Fatal("duplicate packet 100 should be rejected")
}
if !window.Check(50) {
t.Fatal("out-of-order packet 50 within window should be accepted")
}
window.Add(50)
if window.Check(50) {
t.Fatal("duplicate packet 50 should be rejected")
}
// Check packet far behind window (> 8128)
window.Add(10000)
if window.Check(1) {
t.Fatal("packet 1 should be rejected as behind window")
}
}
func TestTCPStream(t *testing.T) {
methods := []struct {
name string
keySize int
}{
{MethodAES128GCM, 16},
{MethodAES256GCM, 32},
{MethodChaCha20Poly1305, 32},
}
dest := net.TCPDestination(net.LocalHostIP, net.Port(8080))
testPayload := []byte("Hello, Shadowsocks 2022 Native Implementation!")
for _, m := range methods {
t.Run(m.name, func(t *testing.T) {
rawKey := make([]byte, m.keySize)
_, _ = rand.Read(rawKey)
method, err := GetCipherMethod(m.name)
common.Must(err)
clientConn, serverConn := gonet.Pipe()
defer clientConn.Close()
defer serverConn.Close()
var wg sync.WaitGroup
wg.Add(2)
var receivedDest net.Destination
var receivedPayload []byte
// Server goroutine
go func() {
defer wg.Done()
salt := make([]byte, method.KeySaltLength)
_, err := io.ReadFull(serverConn, salt)
common.Must(err)
sessionKey := DeriveSessionSubKey(rawKey, salt, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
common.Must(err)
reader := NewStreamReader(serverConn, aead)
// Read fixed chunk (11 + 16 bytes)
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
_, err = io.ReadFull(serverConn, fixedBuf[:])
common.Must(err)
plainFixed, err := aead.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
common.Must(err)
IncreaseNonce(reader.Nonce())
if plainFixed[0] != HeaderTypeClient {
t.Errorf("expected client header type, got %d", plainFixed[0])
}
// Read variable chunk
varLen := int(plainFixed[9])<<8 | int(plainFixed[10])
varBuf := make([]byte, varLen+AEADTagSize)
_, err = io.ReadFull(serverConn, varBuf)
common.Must(err)
plainVar, err := aead.Open(varBuf[:0], reader.Nonce(), varBuf, nil)
common.Must(err)
IncreaseNonce(reader.Nonce())
vBuf := buf.New()
vBuf.Write(plainVar)
receivedDest, err = ReadAddressPort(vBuf)
common.Must(err)
// Skip padding
var padBytes [2]byte
_, _ = vBuf.Read(padBytes[:])
padLen := int(padBytes[0])<<8 | int(padBytes[1])
vBuf.Advance(int32(padLen))
receivedPayload = make([]byte, vBuf.Len())
copy(receivedPayload, vBuf.Bytes())
vBuf.Release()
// Server sends response handshake
serverSalt := make([]byte, method.KeySaltLength)
_, _ = rand.Read(serverSalt)
respKey := DeriveSessionSubKey(rawKey, serverSalt, method.KeySaltLength)
respAead, err := method.NewAEAD(respKey)
writer := NewStreamWriter(serverConn, respAead)
_, _ = serverConn.Write(serverSalt)
fixedResp := make([]byte, 1+8+method.KeySaltLength+2)
fixedResp[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedResp[1:9], uint64(time.Now().Unix()))
copy(fixedResp[9:9+method.KeySaltLength], salt)
binary.BigEndian.PutUint16(fixedResp[9+method.KeySaltLength:11+method.KeySaltLength], 0)
fixedChunk := respAead.Seal(nil, writer.Nonce(), fixedResp, nil)
IncreaseNonce(writer.Nonce())
_, _ = serverConn.Write(fixedChunk)
// Echo stream data
mb, err := reader.ReadMultiBuffer()
common.Must(err)
_ = writer.WriteMultiBuffer(mb)
}()
// Client goroutine
go func() {
defer wg.Done()
clientSalt, writer, err := ClientHandshake(clientConn, method, [][]byte{rawKey}, dest, testPayload)
common.Must(err)
reader, _, err := ClientVerifyServerResponse(clientConn, method, rawKey, clientSalt)
common.Must(err)
// Send additional stream data
streamData := []byte("stream chunk test")
_ = writer.WriteChunk(streamData)
mb, err := reader.ReadMultiBuffer()
common.Must(err)
if !bytes.Equal(mb[0].Bytes(), streamData) {
t.Errorf("echoed stream data mismatch: got %s, want %s", mb[0].Bytes(), streamData)
}
buf.ReleaseMulti(mb)
}()
wg.Wait()
if receivedDest.NetAddr() != dest.NetAddr() {
t.Errorf("destination mismatch: got %s, want %s", receivedDest.NetAddr(), dest.NetAddr())
}
if diff := cmp.Diff(receivedPayload, testPayload); diff != "" {
t.Errorf("payload mismatch: %s", diff)
}
})
}
}
func TestUDPCodec(t *testing.T) {
methods := []string{
MethodAES128GCM,
MethodAES256GCM,
MethodChaCha20Poly1305,
}
dest := net.UDPDestination(net.LocalHostIP, net.Port(53))
payload := []byte("DNS query payload")
for _, methodName := range methods {
t.Run(methodName, func(t *testing.T) {
method, err := GetCipherMethod(methodName)
common.Must(err)
psk := make([]byte, method.KeySaltLength)
_, _ = rand.Read(psk)
clientCodec, err := NewUDPPacketCodec(method, psk)
common.Must(err)
serverCodec, err := NewUDPServerCodec(method, psk, time.Minute)
common.Must(err)
pktBuf, err := clientCodec.EncodeClientPacket(dest, payload)
common.Must(err)
defer pktBuf.Release()
rawCopy := make([]byte, pktBuf.Len())
copy(rawCopy, pktBuf.Bytes())
decoded, err := serverCodec.DecodePacket(pktBuf.Bytes())
common.Must(err)
if decoded.HeaderType != HeaderTypeClient {
t.Errorf("expected header type %d, got %d", HeaderTypeClient, decoded.HeaderType)
}
if decoded.Destination.Port != dest.Port {
t.Errorf("port mismatch: got %d, want %d", decoded.Destination.Port, dest.Port)
}
if !bytes.Equal(decoded.Payload, payload) {
t.Errorf("payload mismatch: got %s, want %s", decoded.Payload, payload)
}
// Replay same packet wire bytes should fail with ErrPacketIdNotUnique
_, err = serverCodec.DecodePacket(rawCopy)
if err != ErrPacketIdNotUnique {
t.Fatalf("expected ErrPacketIdNotUnique on replay, got: %v", err)
}
})
}
}
func TestMultiUserManager(t *testing.T) {
masterKey := generateRandomKey(16)
userKey1 := generateRandomKey(16)
userKey2 := generateRandomKey(16)
config := &MultiUserServerConfig{
Method: MethodAES128GCM,
Key: masterKey,
Users: []*protocol.User{
{
Email: "user1@example.com",
Account: serial.ToTypedMessage(&Account{Key: userKey1}),
},
},
}
inbound, err := NewMultiServer(newTestContext(), config)
common.Must(err)
if inbound.GetUsersCount(context.Background()) != 1 {
t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background()))
}
u1 := inbound.GetUser(context.Background(), "user1@example.com")
if u1 == nil || u1.Email != "user1@example.com" {
t.Fatal("user1 not found")
}
// Add User 2
rawKey2, _ := base64.StdEncoding.DecodeString(userKey2)
u2 := &protocol.MemoryUser{
Email: "user2@example.com",
Account: &MemoryAccount{
Key: rawKey2,
},
}
err = inbound.AddUser(context.Background(), u2)
common.Must(err)
if inbound.GetUsersCount(context.Background()) != 2 {
t.Fatalf("expected 2 users, got %d", inbound.GetUsersCount(context.Background()))
}
// Remove User 1
err = inbound.RemoveUser(context.Background(), "user1@example.com")
common.Must(err)
if inbound.GetUsersCount(context.Background()) != 1 {
t.Fatalf("expected 1 user, got %d", inbound.GetUsersCount(context.Background()))
}
if inbound.GetUser(context.Background(), "user1@example.com") != nil {
t.Fatal("user1 should have been removed")
}
}
+529
View File
@@ -0,0 +1,529 @@
package shadowsocks_2022
import (
"crypto/cipher"
"crypto/rand"
"encoding/binary"
"io"
"math"
mrand "math/rand/v2"
"time"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
)
var addrParser = protocol.NewAddressParser(
protocol.AddressFamilyByte(0x01, net.AddressFamilyIPv4),
protocol.AddressFamilyByte(0x04, net.AddressFamilyIPv6),
protocol.AddressFamilyByte(0x03, net.AddressFamilyDomain),
protocol.WithAddressTypeParser(func(b byte) byte {
return b & 0x0F
}),
)
func IncreaseNonce(nonce []byte) {
for i := range nonce {
nonce[i]++
if nonce[i] != 0 {
return
}
}
}
// WriteAddressPort writes a destination address and port in SOCKS5 format
func WriteAddressPort(w io.Writer, dest net.Destination) error {
return addrParser.WriteAddressPort(w, dest.Address, dest.Port)
}
// ReadAddressPort reads a destination address and port in SOCKS5 format
func ReadAddressPort(r io.Reader) (net.Destination, error) {
addr, port, err := addrParser.ReadAddressPort(nil, r)
if err != nil {
return net.Destination{}, err
}
return net.TCPDestination(addr, port), nil
}
// AddrPortLength returns the serialized length of a destination in SOCKS5 format
func AddrPortLength(dest net.Destination) int {
switch dest.Address.Family() {
case net.AddressFamilyIPv4:
return 1 + 4 + 2
case net.AddressFamilyDomain:
return 1 + 1 + len(dest.Address.Domain()) + 2
case net.AddressFamilyIPv6:
return 1 + 16 + 2
default:
return 0
}
}
type StreamWriter struct {
writer io.Writer
cipher cipher.AEAD
nonce [StreamNonceSize]byte
lenBuf [2]byte
buf []byte
}
func NewStreamWriter(w io.Writer, c cipher.AEAD) *StreamWriter {
return &StreamWriter{
writer: w,
cipher: c,
buf: make([]byte, 0, MaxPacketSize+2+2*AEADTagSize),
}
}
func (w *StreamWriter) Nonce() []byte {
return w.nonce[:]
}
func (w *StreamWriter) WriteChunk(payload []byte) error {
payloadLen := len(payload)
if payloadLen == 0 {
return nil
}
if payloadLen > MaxPacketSize {
return errors.New("payload exceeds MaxPacketSize")
}
binary.BigEndian.PutUint16(w.lenBuf[:], uint16(payloadLen))
w.buf = w.cipher.Seal(w.buf[:0], w.nonce[:], w.lenBuf[:], nil)
IncreaseNonce(w.nonce[:])
w.buf = w.cipher.Seal(w.buf, w.nonce[:], payload, nil)
IncreaseNonce(w.nonce[:])
_, err := w.writer.Write(w.buf)
return err
}
func (w *StreamWriter) Write(p []byte) (int, error) {
n := len(p)
for len(p) > 0 {
chunkSize := len(p)
if chunkSize > MaxPacketSize {
chunkSize = MaxPacketSize
}
if err := w.WriteChunk(p[:chunkSize]); err != nil {
return 0, err
}
p = p[chunkSize:]
}
return n, nil
}
func (w *StreamWriter) WriteMultiBuffer(mb buf.MultiBuffer) error {
defer buf.ReleaseMulti(mb)
for _, b := range mb {
if err := w.WriteChunk(b.Bytes()); err != nil {
return err
}
}
return nil
}
type StreamReader struct {
reader io.Reader
cipher cipher.AEAD
nonce [StreamNonceSize]byte
lenBuf [2 + AEADTagSize]byte
buffer []byte
cached int
offset int
}
func NewStreamReader(r io.Reader, c cipher.AEAD) *StreamReader {
return &StreamReader{
reader: r,
cipher: c,
buffer: make([]byte, MaxPacketSize+AEADTagSize),
}
}
func (r *StreamReader) Nonce() []byte {
return r.nonce[:]
}
func (r *StreamReader) Read(p []byte) (int, error) {
if r.cached > 0 {
n := copy(p, r.buffer[r.offset:r.offset+r.cached])
r.cached -= n
r.offset += n
return n, nil
}
// Read 2-byte length + AEAD tag (18 bytes)
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
return 0, err
}
decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil)
if err != nil {
return 0, errors.New("failed to decrypt chunk length").Base(err)
}
IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 {
return 0, ErrInvalidRequest
}
chunkEnd := payloadLen + AEADTagSize
if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil {
return 0, err
}
decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil)
if err != nil {
return 0, errors.New("failed to decrypt chunk payload").Base(err)
}
IncreaseNonce(r.nonce[:])
r.cached = len(decryptedPayload)
r.offset = 0
n := copy(p, r.buffer[r.offset:r.offset+r.cached])
r.cached -= n
r.offset += n
return n, nil
}
func (r *StreamReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
if r.cached > 0 {
b := buf.New()
b.Write(r.buffer[r.offset : r.offset+r.cached])
r.cached = 0
r.offset = 0
return buf.MultiBuffer{b}, nil
}
if _, err := io.ReadFull(r.reader, r.lenBuf[:]); err != nil {
return nil, err
}
decryptedLen, err := r.cipher.Open(r.lenBuf[:0], r.nonce[:], r.lenBuf[:], nil)
if err != nil {
return nil, errors.New("failed to decrypt chunk length").Base(err)
}
IncreaseNonce(r.nonce[:])
payloadLen := int(binary.BigEndian.Uint16(decryptedLen))
if payloadLen == 0 {
return nil, ErrInvalidRequest
}
chunkEnd := payloadLen + AEADTagSize
if _, err := io.ReadFull(r.reader, r.buffer[:chunkEnd]); err != nil {
return nil, err
}
decryptedPayload, err := r.cipher.Open(r.buffer[:0], r.nonce[:], r.buffer[:chunkEnd], nil)
if err != nil {
return nil, errors.New("failed to decrypt chunk payload").Base(err)
}
IncreaseNonce(r.nonce[:])
b := buf.New()
b.Write(decryptedPayload)
return buf.MultiBuffer{b}, nil
}
type ClientRequestHeader struct {
Destination net.Destination
EarlyData []byte
}
func ReadClientRequestHeader(conn io.Reader, reader *StreamReader) (*ClientRequestHeader, error) {
var fixedBuf [RequestHeaderFixedChunkLength + AEADTagSize]byte
if _, err := io.ReadFull(conn, fixedBuf[:]); err != nil {
return nil, err
}
plainFixed, err := reader.cipher.Open(fixedBuf[:0], reader.Nonce(), fixedBuf[:], nil)
if err != nil {
return nil, errors.New("failed to decrypt client request header").Base(err)
}
IncreaseNonce(reader.Nonce())
if plainFixed[0] != HeaderTypeClient {
return nil, ErrBadHeaderType
}
epoch := binary.BigEndian.Uint64(plainFixed[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(epoch))))
if diff > 30 {
return nil, ErrBadTimestamp
}
varHeaderLen := int(binary.BigEndian.Uint16(plainFixed[9:11]))
if varHeaderLen == 0 {
return nil, ErrInvalidRequest
}
var stackVarChunk [512]byte
var varChunkCipher []byte
needed := varHeaderLen + AEADTagSize
if needed <= len(stackVarChunk) {
varChunkCipher = stackVarChunk[:needed]
} else {
varChunkCipher = make([]byte, needed)
}
if _, err := io.ReadFull(conn, varChunkCipher); err != nil {
return nil, err
}
plainVar, err := reader.cipher.Open(varChunkCipher[:0], reader.Nonce(), varChunkCipher, nil)
if err != nil {
return nil, errors.New("failed to decrypt variable request header").Base(err)
}
IncreaseNonce(reader.Nonce())
b := buf.New()
b.Write(plainVar)
defer b.Release()
dest, err := ReadAddressPort(b)
if err != nil {
return nil, err
}
var padLenBytes [2]byte
if _, err := b.Read(padLenBytes[:]); err != nil {
return nil, err
}
paddingLen := int(binary.BigEndian.Uint16(padLenBytes[:]))
if int(b.Len()) < paddingLen {
return nil, ErrNoPadding
}
if paddingLen > 0 {
b.Advance(int32(paddingLen))
}
var earlyData []byte
if b.Len() > 0 {
earlyData = make([]byte, b.Len())
copy(earlyData, b.Bytes())
}
return &ClientRequestHeader{
Destination: dest,
EarlyData: earlyData,
}, nil
}
// ClientHandshake writes the full client request header to w
func ClientHandshake(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, payload []byte) ([]byte, *StreamWriter, error) {
salt := make([]byte, method.KeySaltLength)
if _, err := io.ReadFull(rand.Reader, salt); err != nil {
return nil, nil, err
}
writer, err := WriteTCPRequest(w, method, pskList, dest, salt, payload)
if err != nil {
return nil, nil, err
}
return salt, writer.(*StreamWriter), nil
}
// ClientVerifyServerResponse reads and verifies the server's handshake response
func ClientVerifyServerResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (*StreamReader, []byte, error) {
reader, err := ReadTCPResponse(r, method, psk, clientSalt)
if err != nil {
return nil, nil, err
}
sr := reader.(*StreamReader)
var initialPayload []byte
if sr.cached > 0 {
initialPayload = make([]byte, sr.cached)
copy(initialPayload, sr.buffer[sr.offset:sr.offset+sr.cached])
}
return sr, initialPayload, nil
}
// WriteTCPRequest writes the Shadowsocks 2022 request header into w and returns a body writer.
func WriteTCPRequest(w io.Writer, method *CipherMethod, pskList [][]byte, dest net.Destination, clientSalt []byte, payload []byte) (buf.Writer, error) {
finalPSK := pskList[len(pskList)-1]
sessionKey := DeriveSessionSubKey(finalPSK, clientSalt, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
return nil, err
}
writer := NewStreamWriter(w, aead)
handshakeBuf := buf.New()
defer handshakeBuf.Release()
handshakeBuf.Write(clientSalt)
for i, currPSK := range pskList[:len(pskList)-1] {
identitySubkey := DeriveIdentitySubKey(currPSK, clientSalt, method.KeySaltLength)
block, err := method.NewBlock(identitySubkey)
if err != nil {
return nil, err
}
nextPSK := pskList[i+1]
pskHash := DeriveUserPSKHash(nextPSK)
var encryptedEIH [AESBlockSize]byte
block.Encrypt(encryptedEIH[:], pskHash[:])
handshakeBuf.Write(encryptedEIH[:])
}
payloadLen := len(payload)
var paddingLen int
if payloadLen < MaxPaddingLength {
paddingLen = mrand.IntN(MaxPaddingLength-payloadLen) + 1
}
addrPortLen := AddrPortLength(dest)
varHeaderLen := addrPortLen + 2 + paddingLen + payloadLen
var fixedHeaderPlaintext [RequestHeaderFixedChunkLength]byte
fixedHeaderPlaintext[0] = HeaderTypeClient
binary.BigEndian.PutUint64(fixedHeaderPlaintext[1:9], uint64(time.Now().Unix()))
binary.BigEndian.PutUint16(fixedHeaderPlaintext[9:11], uint16(varHeaderLen))
fixedChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedHeaderPlaintext[:], nil)
IncreaseNonce(writer.nonce[:])
handshakeBuf.Write(fixedChunk)
varHeaderBuf := buf.New()
defer varHeaderBuf.Release()
if err := WriteAddressPort(varHeaderBuf, dest); err != nil {
return nil, err
}
var padLenBytes [2]byte
binary.BigEndian.PutUint16(padLenBytes[:], uint16(paddingLen))
varHeaderBuf.Write(padLenBytes[:])
if paddingLen > 0 {
varHeaderBuf.Write(zeroPadding[:paddingLen])
}
if payloadLen > 0 {
varHeaderBuf.Write(payload)
}
varChunk := writer.cipher.Seal(nil, writer.nonce[:], varHeaderBuf.Bytes(), nil)
IncreaseNonce(writer.nonce[:])
handshakeBuf.Write(varChunk)
if _, err := w.Write(handshakeBuf.Bytes()); err != nil {
return nil, err
}
return writer, nil
}
// ReadTCPResponse reads and verifies the server's handshake response and returns a reader for the stream.
func ReadTCPResponse(r io.Reader, method *CipherMethod, psk []byte, clientSalt []byte) (buf.Reader, error) {
var serverSalt [32]byte
serverSaltSlice := serverSalt[:method.KeySaltLength]
if _, err := io.ReadFull(r, serverSaltSlice); err != nil {
return nil, err
}
sessionKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
aead, err := method.NewAEAD(sessionKey)
if err != nil {
return nil, err
}
reader := NewStreamReader(r, aead)
fixedPlainLen := 1 + 8 + method.KeySaltLength + 2
chunkCipherLen := fixedPlainLen + AEADTagSize
var chunkBuf [64]byte
chunkSlice := chunkBuf[:chunkCipherLen]
if _, err := io.ReadFull(r, chunkSlice); err != nil {
return nil, err
}
decryptedFixed, err := reader.cipher.Open(chunkSlice[:0], reader.nonce[:], chunkSlice, nil)
if err != nil {
return nil, errors.New("failed to decrypt server response header").Base(err)
}
IncreaseNonce(reader.nonce[:])
if decryptedFixed[0] != HeaderTypeServer {
return nil, ErrBadHeaderType
}
serverEpoch := binary.BigEndian.Uint64(decryptedFixed[1:9])
diff := int(math.Abs(float64(time.Now().Unix() - int64(serverEpoch))))
if diff > 30 {
return nil, ErrBadTimestamp
}
echoedSalt := decryptedFixed[9 : 9+method.KeySaltLength]
for i := 0; i < method.KeySaltLength; i++ {
if echoedSalt[i] != clientSalt[i] {
return nil, errors.New("bad request salt")
}
}
initialPayloadLen := int(binary.BigEndian.Uint16(decryptedFixed[9+method.KeySaltLength : 11+method.KeySaltLength]))
if initialPayloadLen > 0 {
initialCipherLen := initialPayloadLen + AEADTagSize
if _, err := io.ReadFull(r, reader.buffer[:initialCipherLen]); err != nil {
return nil, err
}
decryptedInitial, err := reader.cipher.Open(reader.buffer[:0], reader.nonce[:], reader.buffer[:initialCipherLen], nil)
if err != nil {
return nil, errors.New("failed to decrypt initial response payload").Base(err)
}
IncreaseNonce(reader.nonce[:])
reader.cached = len(decryptedInitial)
reader.offset = 0
}
return reader, nil
}
// WriteTCPResponse writes the server handshake response and returns a body writer for server stream.
func WriteTCPResponse(w io.Writer, method *CipherMethod, psk []byte, clientSalt []byte, initialPayload []byte) (buf.Writer, error) {
var serverSalt [32]byte
serverSaltSlice := serverSalt[:method.KeySaltLength]
if _, err := io.ReadFull(rand.Reader, serverSaltSlice); err != nil {
return nil, err
}
respKey := DeriveSessionSubKey(psk, serverSaltSlice, method.KeySaltLength)
respAead, err := method.NewAEAD(respKey)
if err != nil {
return nil, err
}
writer := NewStreamWriter(w, respAead)
respBuf := buf.New()
defer respBuf.Release()
respBuf.Write(serverSaltSlice)
var fixedRespPlain [1 + 8 + 32 + 2]byte
fixedRespSlice := fixedRespPlain[:1+8+method.KeySaltLength+2]
fixedRespSlice[0] = HeaderTypeServer
binary.BigEndian.PutUint64(fixedRespSlice[1:9], uint64(time.Now().Unix()))
copy(fixedRespSlice[9:9+method.KeySaltLength], clientSalt)
binary.BigEndian.PutUint16(fixedRespSlice[9+method.KeySaltLength:11+method.KeySaltLength], uint16(len(initialPayload)))
fixedRespChunk := writer.cipher.Seal(nil, writer.nonce[:], fixedRespSlice, nil)
IncreaseNonce(writer.nonce[:])
respBuf.Write(fixedRespChunk)
if len(initialPayload) > 0 {
initialChunk := writer.cipher.Seal(nil, writer.nonce[:], initialPayload, nil)
IncreaseNonce(writer.nonce[:])
respBuf.Write(initialChunk)
}
if _, err := w.Write(respBuf.Bytes()); err != nil {
return nil, err
}
return writer, nil
}
+1 -1
View File
@@ -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 {
+2 -2
View File
@@ -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 {
+5 -5
View File
@@ -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
+12 -12
View File
@@ -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

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