diff --git a/adapter/outbound/muxcool.go b/adapter/outbound/muxcool.go new file mode 100644 index 0000000000..59b2cb9e97 --- /dev/null +++ b/adapter/outbound/muxcool.go @@ -0,0 +1,164 @@ +package outbound + +import ( + "context" + "errors" + "fmt" + "net" + "sync" + + "github.com/metacubex/mihomo/common/utils" + C "github.com/metacubex/mihomo/constant" + "github.com/metacubex/mihomo/transport/muxcool" +) + +const ( + muxCoolDestination = "v1.mux.cool" + muxCoolPort = 9527 + + xudpProxyUDP443Reject = "reject" + xudpProxyUDP443Allow = "allow" + xudpProxyUDP443Skip = "skip" +) + +var ErrMuxCoolUDP443Rejected = errors.New("mux.cool rejected UDP/443 traffic") + +type MuxCoolOption struct { + Enabled bool `proxy:"enabled,omitempty"` + MaxConcurrency int `proxy:"max-concurrency,omitempty"` + MaxConnections int `proxy:"max-connections,omitempty"` + MaxCarriers int `proxy:"max-carriers,omitempty"` + XUDPConcurrency int `proxy:"xudp-concurrency,omitempty"` + XUDPProxyUDP443 string `proxy:"xudp-proxy-udp443,omitempty"` +} + +type MuxCool struct { + ProxyAdapter + pool *muxcool.Pool + xudpPool *muxcool.Pool + option MuxCoolOption + closeOnce sync.Once + closeErr error +} + +func NewMuxCool(option MuxCoolOption, proxy ProxyAdapter) (ProxyAdapter, error) { + if option.MaxConcurrency < 0 { + return nil, fmt.Errorf("mux.cool max-concurrency must not be negative") + } + if option.MaxConnections < 0 { + return nil, fmt.Errorf("mux.cool max-connections must not be negative") + } + if option.MaxCarriers < 0 { + return nil, fmt.Errorf("mux.cool max-carriers must not be negative") + } + if option.XUDPConcurrency < 0 { + return nil, fmt.Errorf("mux.cool xudp-concurrency must not be negative") + } + switch option.XUDPProxyUDP443 { + case "": + option.XUDPProxyUDP443 = xudpProxyUDP443Reject + case xudpProxyUDP443Reject, xudpProxyUDP443Allow, xudpProxyUDP443Skip: + default: + return nil, fmt.Errorf("mux.cool xudp-proxy-udp443 must be reject, allow, or skip") + } + if option.MaxConcurrency == 0 { + option.MaxConcurrency = muxcool.DefaultMaxConcurrency + } + if option.MaxConnections == 0 { + option.MaxConnections = muxcool.DefaultMaxConnections + } + + wrapper := &MuxCool{ProxyAdapter: proxy, option: option} + var limiter *muxcool.CarrierLimiter + if option.MaxCarriers > 0 { + limiter = muxcool.NewCarrierLimiter(option.MaxCarriers) + } + wrapper.pool = wrapper.newPool(option.MaxConcurrency, limiter) + if option.XUDPConcurrency > 0 { + wrapper.xudpPool = wrapper.newPool(option.XUDPConcurrency, limiter) + } + return wrapper, nil +} + +func (x *MuxCool) newPool(maxConcurrency int, limiter *muxcool.CarrierLimiter) *muxcool.Pool { + return muxcool.NewPool(x.dialCarrier, muxcool.Options{ + MaxConcurrency: maxConcurrency, + MaxConnections: x.option.MaxConnections, + CarrierLimiter: limiter, + }) +} + +func (x *MuxCool) Options() MuxCoolOption { + return x.option +} + +func (x *MuxCool) dialCarrier(ctx context.Context) (net.Conn, error) { + metadata := &C.Metadata{ + NetWork: C.TCP, + Type: C.INNER, + Host: muxCoolDestination, + DstPort: muxCoolPort, + } + return x.ProxyAdapter.DialContext(ctx, metadata) +} + +func (x *MuxCool) DialContext(ctx context.Context, metadata *C.Metadata) (C.Conn, error) { + if metadata.NetWork != C.TCP { + return x.ProxyAdapter.DialContext(ctx, metadata) + } + conn, err := x.pool.DialContext(ctx, metadata.String(), metadata.DstPort) + if err != nil { + return nil, err + } + return NewConn(conn, x), nil +} + +func (x *MuxCool) ListenPacketContext(ctx context.Context, metadata *C.Metadata) (C.PacketConn, error) { + if metadata.DstPort == 443 { + switch x.option.XUDPProxyUDP443 { + case xudpProxyUDP443Reject: + return nil, ErrMuxCoolUDP443Rejected + case xudpProxyUDP443Skip: + return x.ProxyAdapter.ListenPacketContext(ctx, metadata) + } + } + + var globalID [8]byte + if metadata.SourceValid() { + globalID = utils.GlobalID(metadata.SourceAddress()) + } + pool := x.pool + if x.xudpPool != nil { + pool = x.xudpPool + } + packetConn, err := pool.ListenPacketContext(ctx, metadata.String(), metadata.DstPort, globalID) + if err != nil { + return nil, err + } + return NewPacketConn(packetConn, x), nil +} + +func (x *MuxCool) SupportUDP() bool { + return true +} + +func (x *MuxCool) SupportUOT() bool { + return true +} + +func (x *MuxCool) ProxyInfo() C.ProxyInfo { + info := x.ProxyAdapter.ProxyInfo() + info.XUDP = true + return info +} + +func (x *MuxCool) Close() error { + x.closeOnce.Do(func() { + if x.xudpPool != nil { + _ = x.xudpPool.Close() + } + _ = x.pool.Close() + x.closeErr = x.ProxyAdapter.Close() + }) + return x.closeErr +} diff --git a/adapter/outbound/muxcool_test.go b/adapter/outbound/muxcool_test.go new file mode 100644 index 0000000000..959c4af515 --- /dev/null +++ b/adapter/outbound/muxcool_test.go @@ -0,0 +1,619 @@ +package outbound + +import ( + "context" + "errors" + "net" + "net/netip" + "sync" + "testing" + "time" + + "github.com/metacubex/mihomo/common/utils" + C "github.com/metacubex/mihomo/constant" + "github.com/metacubex/mihomo/transport/muxcool" +) + +type fakeMuxCoolAdapter struct { + *Base + mu sync.Mutex + dialMetadata []C.Metadata + carrierServers chan net.Conn + dialErr error + udpConn C.Conn + packetErr error + packetCalls int + closeCalls int + events []string +} + +func newFakeMuxCoolAdapter() *fakeMuxCoolAdapter { + return &fakeMuxCoolAdapter{ + Base: NewBase(BaseOption{Name: "fake", Addr: "fake:0", Type: C.Direct, UDP: true}), + carrierServers: make(chan net.Conn, 8), + } +} + +func (f *fakeMuxCoolAdapter) DialContext(_ context.Context, metadata *C.Metadata) (C.Conn, error) { + f.mu.Lock() + f.dialMetadata = append(f.dialMetadata, *metadata) + err := f.dialErr + udpConn := f.udpConn + f.mu.Unlock() + if metadata.NetWork == C.UDP && udpConn != nil { + return udpConn, err + } + if err != nil { + return nil, err + } + client, server := net.Pipe() + f.carrierServers <- server + return NewConn(&closeEventConn{Conn: client, onClose: func() { + f.mu.Lock() + f.events = append(f.events, "carrier") + f.mu.Unlock() + }}, f), nil +} + +func (f *fakeMuxCoolAdapter) ListenPacketContext(context.Context, *C.Metadata) (C.PacketConn, error) { + f.mu.Lock() + f.packetCalls++ + err := f.packetErr + f.mu.Unlock() + return nil, err +} + +func (f *fakeMuxCoolAdapter) Close() error { + f.mu.Lock() + f.closeCalls++ + f.events = append(f.events, "base") + f.mu.Unlock() + return nil +} + +type closeEventConn struct { + net.Conn + closeOnce sync.Once + onClose func() +} + +func (c *closeEventConn) Close() error { + c.closeOnce.Do(c.onClose) + return c.Conn.Close() +} + +func (f *fakeMuxCoolAdapter) dialCount() int { + f.mu.Lock() + defer f.mu.Unlock() + return len(f.dialMetadata) +} + +func (f *fakeMuxCoolAdapter) firstMetadata() C.Metadata { + f.mu.Lock() + defer f.mu.Unlock() + return f.dialMetadata[0] +} + +func TestMuxCoolDialsSentinelAndPoolsTCPStreams(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + first, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "one.example", DstPort: 443}) + if err != nil { + t.Fatal(err) + } + second, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "two.example", DstPort: 8443}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = first.Close(); _ = second.Close() }) + + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + writes := []struct { + conn C.Conn + data string + }{ + {conn: first, data: "first"}, + {conn: second, data: "second"}, + } + for _, write := range writes { + done := make(chan error, 1) + go func() { _, err := write.conn.Write([]byte(write.data)); done <- err }() + frame, err := muxcool.DecodeFrame(server) + if err != nil { + t.Fatalf("decode New: %v", err) + } + if frame.Status != muxcool.StatusNew || string(frame.Payload) != write.data { + t.Fatalf("New frame = %+v", frame) + } + if err := <-done; err != nil { + t.Fatalf("logical write: %v", err) + } + } + if got := base.dialCount(); got != 1 { + t.Fatalf("base dial count = %d, want 1", got) + } + carrierMetadata := base.firstMetadata() + if carrierMetadata.NetWork != C.TCP || carrierMetadata.Host != muxCoolDestination || carrierMetadata.DstPort != muxCoolPort { + t.Fatalf("carrier metadata = %+v", carrierMetadata) + } +} + +func TestMuxCoolDoesNotBypassFailedCarrier(t *testing.T) { + dialErr := errors.New("sentinel unavailable") + base := newFakeMuxCoolAdapter() + base.dialErr = dialErr + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + conn, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "target.example", DstPort: 443}) + if conn != nil || !errors.Is(err, dialErr) { + t.Fatalf("DialContext = (%v, %v), want carrier error", conn, err) + } + if got := base.dialCount(); got != 1 { + t.Fatalf("base dial count = %d, want no fallback", got) + } + if metadata := base.firstMetadata(); metadata.Host != muxCoolDestination { + t.Fatalf("dialed host = %q", metadata.Host) + } +} + +func TestMuxCoolDialContextDelegatesUDPStreams(t *testing.T) { + base := newFakeMuxCoolAdapter() + udpClient, udpPeer := net.Pipe() + t.Cleanup(func() { _ = udpPeer.Close() }) + base.udpConn = NewConn(udpClient, base) + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + metadata := &C.Metadata{NetWork: C.UDP, Host: "udp.example", DstPort: 53} + + conn, err := wrapped.DialContext(context.Background(), metadata) + if err != nil || conn != base.udpConn { + t.Fatalf("UDP DialContext = (%v, %v)", conn, err) + } +} + +func TestMuxCoolPoolsUDPAndDerivesXUDPGlobalID(t *testing.T) { + base := newFakeMuxCoolAdapter() + base.Base.udp = false + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + metadata := &C.Metadata{ + NetWork: C.UDP, + Host: "dns.example", + DstPort: 53, + SrcIP: netip.MustParseAddr("192.0.2.10"), + SrcPort: 4242, + } + + packetConn, err := wrapped.ListenPacketContext(context.Background(), metadata) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packetConn.Close() }) + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + + queryAddr := net.UDPAddrFromAddrPort(netip.MustParseAddrPort("8.8.8.8:53")) + writeDone := make(chan error, 1) + go func() { + _, err := packetConn.WriteTo([]byte("query"), queryAddr) + writeDone <- err + }() + frame, err := muxcool.DecodeFrame(server) + if err != nil { + t.Fatal(err) + } + if err := <-writeDone; err != nil { + t.Fatal(err) + } + if frame.Status != muxcool.StatusNew || frame.Network != muxcool.NetworkUDP || frame.Destination != "dns.example" || frame.Port != 53 { + t.Fatalf("UDP New = %+v", frame) + } + if want := utils.GlobalID(metadata.SourceAddress()); frame.GlobalID != want { + t.Fatalf("GlobalID = %v, want %v", frame.GlobalID, want) + } + if string(frame.Payload) != "query" { + t.Fatalf("payload = %q", frame.Payload) + } + + response, err := muxcool.EncodeFrame(muxcool.Frame{ + SessionID: frame.SessionID, Status: muxcool.StatusKeep, Option: muxcool.OptionData, + Network: muxcool.NetworkUDP, Destination: "8.8.4.4", Port: 53, Payload: []byte("answer"), + }) + if err != nil { + t.Fatal(err) + } + go func() { _, _ = server.Write(response) }() + buffer := make([]byte, 16) + n, addr, err := packetConn.ReadFrom(buffer) + if err != nil { + t.Fatal(err) + } + if string(buffer[:n]) != "answer" || addr.String() != "8.8.4.4:53" { + t.Fatalf("ReadFrom = (%q, %v)", buffer[:n], addr) + } + + base.mu.Lock() + packetCalls := base.packetCalls + base.mu.Unlock() + if packetCalls != 0 { + t.Fatalf("base packet calls = %d, want 0", packetCalls) + } + if !wrapped.SupportUDP() || !wrapped.SupportUOT() || !wrapped.ProxyInfo().XUDP { + t.Fatalf("capabilities = UDP %t UOT %t XUDP %t", wrapped.SupportUDP(), wrapped.SupportUOT(), wrapped.ProxyInfo().XUDP) + } +} + +func TestMuxCoolUsesNormalUDPWithoutSourceIdentity(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + packetConn, err := wrapped.ListenPacketContext(context.Background(), &C.Metadata{NetWork: C.UDP, Host: "dns.example", DstPort: 53}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packetConn.Close() }) + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + + go func() { + _, _ = packetConn.WriteTo([]byte("query"), net.UDPAddrFromAddrPort(netip.MustParseAddrPort("1.1.1.1:53"))) + }() + frame, err := muxcool.DecodeFrame(server) + if err != nil { + t.Fatal(err) + } + if frame.GlobalID != [8]byte{} { + t.Fatalf("GlobalID = %v, want normal UDP", frame.GlobalID) + } +} + +func TestMuxCoolUsesDedicatedPoolWhenXUDPConcurrencyIsPositive(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{ + Enabled: true, + MaxConcurrency: 8, + XUDPConcurrency: 4, + XUDPProxyUDP443: "allow", + }, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + stream, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "tcp.example", DstPort: 80}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = stream.Close() }) + packetConn, err := wrapped.ListenPacketContext(context.Background(), &C.Metadata{NetWork: C.UDP, Host: "udp.example", DstPort: 53}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packetConn.Close() }) + + if got := base.dialCount(); got != 2 { + t.Fatalf("carrier dial count = %d, want separate TCP and XUDP carriers", got) + } + for i := 0; i < 2; i++ { + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + } +} + +func TestMuxCoolSharesMaxCarriersAcrossTCPAndXUDPPools(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{ + Enabled: true, + MaxConcurrency: 1, + MaxCarriers: 1, + XUDPConcurrency: 1, + }, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + stream, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "tcp.example", DstPort: 80}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = stream.Close() }) + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + packetConn, err := wrapped.ListenPacketContext(ctx, &C.Metadata{NetWork: C.UDP, Host: "udp.example", DstPort: 53}) + if packetConn != nil || !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("XUDP at shared carrier limit = (%v, %v), want deadline exceeded", packetConn, err) + } + if got := base.dialCount(); got != 1 { + t.Fatalf("carrier dial count = %d, want strict shared limit of 1", got) + } +} + +func TestMuxCoolSharesPoolWhenXUDPConcurrencyIsZero(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + stream, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "tcp.example", DstPort: 80}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = stream.Close() }) + packetConn, err := wrapped.ListenPacketContext(context.Background(), &C.Metadata{NetWork: C.UDP, Host: "udp.example", DstPort: 53}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packetConn.Close() }) + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + + if got := base.dialCount(); got != 1 { + t.Fatalf("carrier dial count = %d, want shared TCP and XUDP carrier", got) + } +} + +func TestMuxCoolEnforcesDedicatedXUDPConcurrency(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true, XUDPConcurrency: 1}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + var packetConns []C.PacketConn + for i := 0; i < 2; i++ { + packetConn, err := wrapped.ListenPacketContext(context.Background(), &C.Metadata{NetWork: C.UDP, Host: "udp.example", DstPort: 53}) + if err != nil { + t.Fatal(err) + } + packetConns = append(packetConns, packetConn) + } + t.Cleanup(func() { + for _, packetConn := range packetConns { + _ = packetConn.Close() + } + }) + + if got := base.dialCount(); got != 2 { + t.Fatalf("carrier dial count = %d, want 2 at XUDP concurrency 1", got) + } + for i := 0; i < 2; i++ { + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + } +} + +func TestMuxCoolRejectsUDP443ByDefault(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + packetConn, err := wrapped.ListenPacketContext(context.Background(), &C.Metadata{NetWork: C.UDP, Host: "quic.example", DstPort: 443}) + if packetConn != nil || !errors.Is(err, ErrMuxCoolUDP443Rejected) { + t.Fatalf("ListenPacketContext = (%v, %v), want UDP/443 rejection", packetConn, err) + } + if got := base.dialCount(); got != 0 { + t.Fatalf("carrier dial count = %d, want 0", got) + } +} + +func TestMuxCoolAllowsUDP443(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true, XUDPProxyUDP443: "allow"}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + packetConn, err := wrapped.ListenPacketContext(context.Background(), &C.Metadata{NetWork: C.UDP, Host: "quic.example", DstPort: 443}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packetConn.Close() }) + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + + if got := base.dialCount(); got != 1 { + t.Fatalf("carrier dial count = %d, want 1", got) + } + base.mu.Lock() + packetCalls := base.packetCalls + base.mu.Unlock() + if packetCalls != 0 { + t.Fatalf("base packet calls = %d, want 0", packetCalls) + } +} + +func TestMuxCoolSkipsUDP443(t *testing.T) { + skipErr := errors.New("base UDP called") + base := newFakeMuxCoolAdapter() + base.packetErr = skipErr + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true, XUDPProxyUDP443: "skip"}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + + packetConn, err := wrapped.ListenPacketContext(context.Background(), &C.Metadata{NetWork: C.UDP, Host: "quic.example", DstPort: 443}) + if packetConn != nil || !errors.Is(err, skipErr) { + t.Fatalf("ListenPacketContext = (%v, %v), want base proxy result", packetConn, err) + } + if got := base.dialCount(); got != 0 { + t.Fatalf("carrier dial count = %d, want 0", got) + } + base.mu.Lock() + packetCalls := base.packetCalls + base.mu.Unlock() + if packetCalls != 1 { + t.Fatalf("base packet calls = %d, want 1", packetCalls) + } +} + +func TestMuxCoolServerFirstNewFrame(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + conn, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "smtp.example", DstPort: 25}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = conn.Close() }) + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + _ = server.SetReadDeadline(time.Now().Add(time.Second)) + frame, err := muxcool.DecodeFrame(server) + if err != nil { + t.Fatal(err) + } + if frame.Status != muxcool.StatusNew || frame.Option&muxcool.OptionData != 0 || len(frame.Payload) != 0 { + t.Fatalf("server-first frame = %+v", frame) + } +} + +func TestMuxCoolContextCancellationClosesOnlyLogicalSession(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = wrapped.Close() }) + ctx, cancel := context.WithCancel(context.Background()) + first, err := wrapped.DialContext(ctx, &C.Metadata{NetWork: C.TCP, Host: "cancel.example", DstPort: 80}) + if err != nil { + t.Fatal(err) + } + second, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "survivor.example", DstPort: 80}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = first.Close(); _ = second.Close() }) + server := <-base.carrierServers + t.Cleanup(func() { _ = server.Close() }) + + cancel() + end, err := muxcool.DecodeFrame(server) + if err != nil { + t.Fatalf("decode cancelled End: %v", err) + } + if end.SessionID != 1 || end.Status != muxcool.StatusEnd { + t.Fatalf("cancel frame = %+v", end) + } + if _, err := first.Read(make([]byte, 1)); !errors.Is(err, context.Canceled) { + t.Fatalf("cancelled read error = %v", err) + } + + response, err := muxcool.EncodeFrame(muxcool.Frame{SessionID: 2, Status: muxcool.StatusKeep, Option: muxcool.OptionData, Payload: []byte("ok")}) + if err != nil { + t.Fatal(err) + } + go func() { _, _ = server.Write(response) }() + got := make([]byte, 2) + if _, err := second.Read(got); err != nil { + t.Fatalf("surviving session read: %v", err) + } + if string(got) != "ok" { + t.Fatalf("surviving response = %q", got) + } +} + +func TestMuxCoolCloseClosesPoolBeforeBaseExactlyOnce(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true}, base) + if err != nil { + t.Fatal(err) + } + conn, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "close.example", DstPort: 80}) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + server := <-base.carrierServers + defer server.Close() + + if err := wrapped.Close(); err != nil { + t.Fatal(err) + } + if err := wrapped.Close(); err != nil { + t.Fatal(err) + } + base.mu.Lock() + events := append([]string(nil), base.events...) + closeCalls := base.closeCalls + base.mu.Unlock() + if closeCalls != 1 { + t.Fatalf("base close calls = %d", closeCalls) + } + if len(events) < 2 || events[0] != "carrier" || events[len(events)-1] != "base" { + t.Fatalf("close events = %v, want carrier before base", events) + } +} + +func TestMuxCoolCloseClosesDedicatedXUDPPoolBeforeBase(t *testing.T) { + base := newFakeMuxCoolAdapter() + wrapped, err := NewMuxCool(MuxCoolOption{Enabled: true, XUDPConcurrency: 4}, base) + if err != nil { + t.Fatal(err) + } + stream, err := wrapped.DialContext(context.Background(), &C.Metadata{NetWork: C.TCP, Host: "close.example", DstPort: 80}) + if err != nil { + t.Fatal(err) + } + defer stream.Close() + packetConn, err := wrapped.ListenPacketContext(context.Background(), &C.Metadata{NetWork: C.UDP, Host: "close.example", DstPort: 53}) + if err != nil { + t.Fatal(err) + } + defer packetConn.Close() + for i := 0; i < 2; i++ { + server := <-base.carrierServers + defer server.Close() + } + + if err := wrapped.Close(); err != nil { + t.Fatal(err) + } + if err := wrapped.Close(); err != nil { + t.Fatal(err) + } + base.mu.Lock() + events := append([]string(nil), base.events...) + closeCalls := base.closeCalls + base.mu.Unlock() + if closeCalls != 1 { + t.Fatalf("base close calls = %d, want 1", closeCalls) + } + if len(events) != 3 || events[0] != "carrier" || events[1] != "carrier" || events[2] != "base" { + t.Fatalf("close events = %v, want both carriers before base", events) + } +} diff --git a/adapter/outbound/vless.go b/adapter/outbound/vless.go index 4d076b5925..517d707f41 100644 --- a/adapter/outbound/vless.go +++ b/adapter/outbound/vless.go @@ -431,6 +431,9 @@ func (v *Vless) Close() error { } func parseVlessAddr(metadata *C.Metadata, xudp bool) *vless.DstAddr { + nativeMux := metadata.NetWork == C.TCP && + metadata.Host == muxCoolDestination && + metadata.DstPort == muxCoolPort var addrType byte var addr []byte switch metadata.AddrType() { @@ -454,7 +457,7 @@ func parseVlessAddr(metadata *C.Metadata, xudp bool) *vless.DstAddr { AddrType: addrType, Addr: addr, Port: metadata.DstPort, - Mux: metadata.NetWork == C.UDP && xudp, + Mux: nativeMux || metadata.NetWork == C.UDP && xudp, } } diff --git a/adapter/outbound/vless_test.go b/adapter/outbound/vless_test.go new file mode 100644 index 0000000000..2fdfa473dc --- /dev/null +++ b/adapter/outbound/vless_test.go @@ -0,0 +1,34 @@ +package outbound + +import ( + "testing" + + C "github.com/metacubex/mihomo/constant" +) + +func TestParseVlessAddrUsesNativeMuxCommandForMuxCoolCarrier(t *testing.T) { + metadata := &C.Metadata{ + NetWork: C.TCP, + Type: C.INNER, + Host: muxCoolDestination, + DstPort: muxCoolPort, + } + + destination := parseVlessAddr(metadata, false) + if !destination.Mux { + t.Fatal("mux.cool carrier did not select the native VLESS Mux command") + } +} + +func TestParseVlessAddrDoesNotUseMuxCommandForOrdinaryTCP(t *testing.T) { + metadata := &C.Metadata{ + NetWork: C.TCP, + Host: muxCoolDestination, + DstPort: muxCoolPort + 1, + } + + destination := parseVlessAddr(metadata, false) + if destination.Mux { + t.Fatal("ordinary TCP destination selected the native VLESS Mux command") + } +} diff --git a/adapter/parser.go b/adapter/parser.go index 9599616a19..8fa74729b5 100644 --- a/adapter/parser.go +++ b/adapter/parser.go @@ -217,17 +217,34 @@ func ParseProxy(mapping map[string]any, options ...ProxyOption) (C.Proxy, error) return nil, err } + muxOption := &outbound.SingMuxOption{} if muxMapping, muxExist := mapping["smux"].(map[string]any); muxExist { - muxOption := &outbound.SingMuxOption{} err = decoder.Decode(muxMapping, muxOption) if err != nil { return nil, err } - if muxOption.Enabled { - proxy, err = outbound.NewSingMux(*muxOption, proxy) - if err != nil { - return nil, err - } + } + + muxCoolOption := &outbound.MuxCoolOption{} + if muxMapping, muxExist := mapping["mux.cool"].(map[string]any); muxExist { + err = decoder.Decode(muxMapping, muxCoolOption) + if err != nil { + return nil, err + } + } + if muxOption.Enabled && muxCoolOption.Enabled { + return nil, fmt.Errorf("smux and mux.cool cannot be enabled together") + } + if muxOption.Enabled { + proxy, err = outbound.NewSingMux(*muxOption, proxy) + if err != nil { + return nil, err + } + } + if muxCoolOption.Enabled || muxCoolOption.MaxConcurrency < 0 || muxCoolOption.MaxConnections < 0 || muxCoolOption.MaxCarriers < 0 || muxCoolOption.XUDPConcurrency < 0 { + proxy, err = outbound.NewMuxCool(*muxCoolOption, proxy) + if err != nil { + return nil, err } } diff --git a/adapter/parser_muxcool_test.go b/adapter/parser_muxcool_test.go new file mode 100644 index 0000000000..97d06d453d --- /dev/null +++ b/adapter/parser_muxcool_test.go @@ -0,0 +1,135 @@ +package adapter + +import ( + "reflect" + "strings" + "testing" + + "github.com/metacubex/mihomo/adapter/outbound" +) + +func TestParseProxyMuxCoolConfigurations(t *testing.T) { + tests := []struct { + name string + muxCool map[string]any + wantWrapped bool + wantMaxConcurrency int + wantMaxConnections int + wantMaxCarriers int + wantXUDPConcurrency int + wantXUDPProxyUDP443Mode string + }{ + {name: "omitted"}, + {name: "disabled", muxCool: map[string]any{"enabled": false}}, + { + name: "defaults", muxCool: map[string]any{"enabled": true}, wantWrapped: true, + wantMaxConcurrency: 8, wantMaxConnections: 128, wantXUDPProxyUDP443Mode: "reject", + }, + { + name: "custom", + muxCool: map[string]any{ + "enabled": true, "max-concurrency": 3, "max-connections": 19, "max-carriers": 2, + "xudp-concurrency": 4, "xudp-proxy-udp443": "allow", + }, + wantWrapped: true, wantMaxConcurrency: 3, wantMaxConnections: 19, + wantMaxCarriers: 2, wantXUDPConcurrency: 4, wantXUDPProxyUDP443Mode: "allow", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + mapping := map[string]any{"name": "test", "type": "direct"} + if tt.muxCool != nil { + mapping["mux.cool"] = tt.muxCool + } + parsed, err := ParseProxy(mapping) + if err != nil { + t.Fatalf("ParseProxy: %v", err) + } + defer parsed.Close() + wrapped := unwrapAutoClose(t, parsed.Adapter()) + muxCool, ok := wrapped.(*outbound.MuxCool) + if ok != tt.wantWrapped { + t.Fatalf("MuxCool wrapper present = %t, want %t (%T)", ok, tt.wantWrapped, wrapped) + } + if ok { + options := muxCool.Options() + if options.MaxConcurrency != tt.wantMaxConcurrency || + options.MaxConnections != tt.wantMaxConnections || + options.MaxCarriers != tt.wantMaxCarriers || + options.XUDPConcurrency != tt.wantXUDPConcurrency || + options.XUDPProxyUDP443 != tt.wantXUDPProxyUDP443Mode { + t.Fatalf("options = %+v", options) + } + } + }) + } +} + +func TestParseProxyRejectsInvalidMuxCoolConfigurations(t *testing.T) { + tests := []struct { + name string + mapping map[string]any + match string + }{ + { + name: "negative concurrency", + mapping: map[string]any{"name": "test", "type": "direct", "mux.cool": map[string]any{"enabled": true, "max-concurrency": -1}}, + match: "max-concurrency", + }, + { + name: "negative connections", + mapping: map[string]any{"name": "test", "type": "direct", "mux.cool": map[string]any{"enabled": true, "max-connections": -1}}, + match: "max-connections", + }, + { + name: "negative carriers", + mapping: map[string]any{"name": "test", "type": "direct", "mux.cool": map[string]any{"enabled": true, "max-carriers": -1}}, + match: "max-carriers", + }, + { + name: "negative xudp concurrency", + mapping: map[string]any{"name": "test", "type": "direct", "mux.cool": map[string]any{"enabled": true, "xudp-concurrency": -1}}, + match: "xudp-concurrency", + }, + { + name: "unknown udp 443 policy", + mapping: map[string]any{"name": "test", "type": "direct", "mux.cool": map[string]any{"enabled": true, "xudp-proxy-udp443": "drop"}}, + match: "xudp-proxy-udp443", + }, + { + name: "smux conflict", + mapping: map[string]any{ + "name": "test", "type": "direct", + "smux": map[string]any{"enabled": true}, + "mux.cool": map[string]any{"enabled": true}, + }, + match: "cannot be enabled together", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + parsed, err := ParseProxy(tt.mapping) + if parsed != nil { + _ = parsed.Close() + } + if err == nil || !strings.Contains(err.Error(), tt.match) { + t.Fatalf("error = %v, want substring %q", err, tt.match) + } + }) + } +} + +func unwrapAutoClose(t *testing.T, adapter any) any { + t.Helper() + value := reflect.ValueOf(adapter) + if value.Kind() != reflect.Pointer || value.Elem().Kind() != reflect.Struct { + t.Fatalf("unexpected outer adapter %T", adapter) + } + field := value.Elem().FieldByName("ProxyAdapter") + if !field.IsValid() || !field.CanInterface() { + t.Fatalf("cannot unwrap outer adapter %T", adapter) + } + return field.Interface() +} diff --git a/docs/config.yaml b/docs/config.yaml index 80077927b0..2736ae260c 100644 --- a/docs/config.yaml +++ b/docs/config.yaml @@ -486,6 +486,16 @@ proxies: # socks5 # statistic: false # 控制是否将底层连接显示在面板中,方便打断底层连接 # only-tcp: false # 如果设置为 true, smux 的设置将不会对 udp 生效,udp 连接会直接走底层协议 + # mux.cool is Xray-compatible multiplexing and is separate from smux above. Do not enable both. + # It carries TCP streams and UDP/XUDP packet sessions over shared TCP carriers. + # Packet loss can cause head-of-line blocking for every logical session on the carrier. + mux.cool: + enabled: false + # max-concurrency: 8 # Maximum active TCP/UDP sessions per shared carrier. 0 uses the default (8). + # max-connections: 128 # Maximum lifetime sessions per carrier before rotation. 0 uses the default (128). + # xudp-concurrency: 4 # Dedicated XUDP sessions per carrier. 0 shares the main pool. + # xudp-proxy-udp443: reject # reject, allow, or skip mux.cool for UDP/443. + - name: "ss2" type: ss server: server diff --git a/listener/inbound/mux_test.go b/listener/inbound/mux_test.go index 96841dbf1b..d2c90b5c4f 100644 --- a/listener/inbound/mux_test.go +++ b/listener/inbound/mux_test.go @@ -13,12 +13,28 @@ var singMuxProtocolList = []string{"h2mux", "smux", "yamux"} var singMuxProtocolListLong = []string{"yamux"} // don't test "smux", "h2mux" because it has some confused bugs // notCloseProxyAdapter is a proxy adapter that does not close the underlying outbound.ProxyAdapter. -// The outbound.SingMux will close the underlying outbound.ProxyAdapter when it is closed, but we don't want to close it. +// Multiplexing wrappers close their underlying ProxyAdapter, but the test owner keeps it alive. // The underlying outbound.ProxyAdapter should only be closed by the caller of testSingMux. type notCloseProxyAdapter struct { outbound.ProxyAdapter } +func testMuxCool(t *testing.T, tunnel *TestTunnel, out outbound.ProxyAdapter) { + t.Run("mux.cool", func(t *testing.T) { + muxCool, err := outbound.NewMuxCool(outbound.MuxCoolOption{ + Enabled: true, + XUDPProxyUDP443: "allow", + }, ¬CloseProxyAdapter{out}) + if !assert.NoError(t, err) { + return + } + defer muxCool.Close() + + tunnel.DoSequentialTest(t, muxCool) + tunnel.DoConcurrentTest(t, muxCool) + }) +} + func (n *notCloseProxyAdapter) Close() error { return nil } diff --git a/listener/inbound/vless_test.go b/listener/inbound/vless_test.go index d83df33e55..3a6ba5ca94 100644 --- a/listener/inbound/vless_test.go +++ b/listener/inbound/vless_test.go @@ -12,6 +12,10 @@ import ( ) func testInboundVless(t *testing.T, inboundOptions inbound.VlessOption, outboundOptions outbound.VlessOption) { + testInboundVlessWithMuxCool(t, inboundOptions, outboundOptions, false) +} + +func testInboundVlessWithMuxCool(t *testing.T, inboundOptions inbound.VlessOption, outboundOptions outbound.VlessOption, runMuxCool bool) { t.Parallel() inboundOptions.BaseOption = inbound.BaseOption{ NameStr: "vless_inbound", @@ -59,6 +63,13 @@ func testInboundVless(t *testing.T, inboundOptions inbound.VlessOption, outbound return } testSingMux(t, tunnel, out) + if runMuxCool { + testMuxCool(t, tunnel, out) + } +} + +func TestInboundVless_MuxCool(t *testing.T) { + testInboundVlessWithMuxCool(t, inbound.VlessOption{AllowInsecure: true}, outbound.VlessOption{}, true) } func testInboundVlessTLS(t *testing.T, inboundOptions inbound.VlessOption, outboundOptions outbound.VlessOption, testVision bool) { diff --git a/listener/inbound/vmess_test.go b/listener/inbound/vmess_test.go index af4a8bd568..057f0d54f5 100644 --- a/listener/inbound/vmess_test.go +++ b/listener/inbound/vmess_test.go @@ -13,6 +13,10 @@ import ( ) func testInboundVMess(t *testing.T, inboundOptions inbound.VmessOption, outboundOptions outbound.VmessOption) { + testInboundVMessWithMuxCool(t, inboundOptions, outboundOptions, false) +} + +func testInboundVMessWithMuxCool(t *testing.T, inboundOptions inbound.VmessOption, outboundOptions outbound.VmessOption, runMuxCool bool) { t.Parallel() inboundOptions.BaseOption = inbound.BaseOption{ NameStr: "vmess_inbound", @@ -71,6 +75,9 @@ func testInboundVMess(t *testing.T, inboundOptions inbound.VmessOption, outbound return } testSingMux(t, tunnel, out) + if runMuxCool { + testMuxCool(t, tunnel, out) + } } func TestInboundVMess_Basic(t *testing.T) { @@ -79,6 +86,10 @@ func TestInboundVMess_Basic(t *testing.T) { testInboundVMess(t, inboundOptions, outboundOptions) } +func TestInboundVMess_MuxCool(t *testing.T) { + testInboundVMessWithMuxCool(t, inbound.VmessOption{}, outbound.VmessOption{}, true) +} + func testInboundVMessTLS(t *testing.T, inboundOptions inbound.VmessOption, outboundOptions outbound.VmessOption) { testInboundVMess(t, inboundOptions, outboundOptions) t.Run("ECH", func(t *testing.T) { diff --git a/listener/sing/muxcool_test.go b/listener/sing/muxcool_test.go new file mode 100644 index 0000000000..a74e1466b4 --- /dev/null +++ b/listener/sing/muxcool_test.go @@ -0,0 +1,50 @@ +package sing + +import ( + "context" + "errors" + "io" + "net" + "testing" + "time" + + "github.com/metacubex/mihomo/transport/muxcool" + + vmess "github.com/metacubex/sing-vmess" + M "github.com/metacubex/sing/common/metadata" +) + +func TestListenerHandlerCloseDrainsMuxCoolCarrier(t *testing.T) { + handler, err := NewListenerHandler(ListenerConfig{}) + if err != nil { + t.Fatal(err) + } + client, server := net.Pipe() + done := make(chan error, 1) + go func() { + done <- handler.ParseSpecialFqdn(context.Background(), server, M.Metadata{ + Source: M.ParseSocksaddr("192.0.2.20:32000"), + Destination: vmess.MuxDestination, + }) + }() + keepAlive, err := muxcool.EncodeFrame(muxcool.Frame{Status: muxcool.StatusKeepAlive}) + if err != nil { + t.Fatal(err) + } + if _, err := client.Write(keepAlive); err != nil { + t.Fatal(err) + } + + if err := handler.Close(); err != nil { + t.Fatal(err) + } + select { + case err := <-done: + if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, io.ErrClosedPipe) && !errors.Is(err, net.ErrClosed) { + t.Fatalf("ParseSpecialFqdn: %v", err) + } + case <-time.After(time.Second): + t.Fatal("mux.cool carrier was not drained") + } + _ = client.Close() +} diff --git a/listener/sing/sing.go b/listener/sing/sing.go index 7bb741734d..f2df7708db 100644 --- a/listener/sing/sing.go +++ b/listener/sing/sing.go @@ -14,6 +14,7 @@ import ( "github.com/metacubex/mihomo/common/utils" C "github.com/metacubex/mihomo/constant" "github.com/metacubex/mihomo/log" + "github.com/metacubex/mihomo/transport/muxcool" "github.com/gofrs/uuid/v5" mux "github.com/metacubex/sing-mux" @@ -55,6 +56,7 @@ type ListenerHandler struct { ListenerConfig handlerId uuid.UUID muxService *mux.Service + muxCool *muxcool.ServerRuntime } func UpstreamMetadata(metadata M.Metadata) M.Metadata { @@ -75,6 +77,7 @@ func ConvertMetadata(metadata *C.Metadata) M.Metadata { func NewListenerHandler(lc ListenerConfig) (h *ListenerHandler, err error) { h = &ListenerHandler{ListenerConfig: lc} h.handlerId = utils.NewUUIDV4() + h.muxCool = muxcool.NewServerRuntime(muxcool.ServerOptions{}) h.muxService, err = mux.NewService(mux.ServiceOptions{ NewStreamContext: func(ctx context.Context, conn net.Conn) context.Context { return ctx @@ -108,7 +111,7 @@ func (h *ListenerHandler) ParseSpecialFqdn(ctx context.Context, conn net.Conn, m case mux.Destination.Fqdn: return h.muxService.NewConnection(ctx, conn, UpstreamMetadata(metadata)) case vmess.MuxDestination.Fqdn: - return vmess.HandleMuxConnection(ctx, conn, metadata, h) + return h.muxCool.Serve(ctx, conn, metadata, h) case uot.MagicAddress: request, err := uot.ReadRequest(conn) if err != nil { @@ -123,6 +126,10 @@ func (h *ListenerHandler) ParseSpecialFqdn(ctx context.Context, conn net.Conn, m return errors.New("not special fqdn") } +func (h *ListenerHandler) Close() error { + return h.muxCool.Close() +} + func (h *ListenerHandler) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error { if h.IsSpecialFqdn(metadata.Destination.Fqdn) { return h.ParseSpecialFqdn(ctx, conn, metadata) diff --git a/listener/sing_vless/server.go b/listener/sing_vless/server.go index e1237ddc88..a6c1c86400 100644 --- a/listener/sing_vless/server.go +++ b/listener/sing_vless/server.go @@ -36,6 +36,7 @@ type Listener struct { listeners []net.Listener service *Service[string] decryption *encryption.ServerInstance + handler *sing.ListenerHandler } func New(config LC.VlessServer, lc C.InboundListenConfig, tunnel C.Tunnel, additions ...inbound.Addition) (sl *Listener, err error) { @@ -67,7 +68,7 @@ func New(config LC.VlessServer, lc C.InboundListenConfig, tunnel C.Tunnel, addit return it.Flow })) - sl = &Listener{config: config, service: service} + sl = &Listener{config: config, service: service, handler: h} sl.decryption, err = encryption.NewServer(config.Decryption) if err != nil { @@ -311,6 +312,9 @@ func (l *Listener) Close() error { if l.decryption != nil { _ = l.decryption.Close() } + if err := l.handler.Close(); err != nil { + retErr = err + } return retErr } diff --git a/listener/sing_vmess/server.go b/listener/sing_vmess/server.go index 156432b9e2..28cd04d3ae 100644 --- a/listener/sing_vmess/server.go +++ b/listener/sing_vmess/server.go @@ -38,6 +38,7 @@ type Listener struct { config LC.VmessServer listeners []net.Listener service *vmess.Service[string] + handler *sing.ListenerHandler } var _listener *Listener @@ -90,7 +91,7 @@ func New(config LC.VmessServer, lc C.InboundListenConfig, tunnel C.Tunnel, addit return nil, err } - sl = &Listener{false, config, nil, service} + sl = &Listener{closed: false, config: config, service: service, handler: h} httpServer := http.Server{ IdleTimeout: 30 * time.Second, @@ -317,6 +318,9 @@ func (l *Listener) Close() error { if err != nil { retErr = err } + if err := l.handler.Close(); err != nil { + retErr = err + } return retErr } diff --git a/transport/muxcool/carrier_limiter.go b/transport/muxcool/carrier_limiter.go new file mode 100644 index 0000000000..41394ce469 --- /dev/null +++ b/transport/muxcool/carrier_limiter.go @@ -0,0 +1,68 @@ +package muxcool + +import ( + "net" + "sync" +) + +// CarrierLimiter enforces a shared upper bound on physical mux.cool carriers. +// A zero limit preserves the existing unlimited behavior. +type CarrierLimiter struct { + mu sync.Mutex + max int + used int + changed chan struct{} +} + +func NewCarrierLimiter(maxCarriers int) *CarrierLimiter { + return &CarrierLimiter{ + max: maxCarriers, + changed: make(chan struct{}), + } +} + +func (l *CarrierLimiter) limited() bool { + return l != nil && l.max > 0 +} + +func (l *CarrierLimiter) tryAcquire() (*carrierLease, <-chan struct{}) { + l.mu.Lock() + defer l.mu.Unlock() + if l.used >= l.max { + return nil, l.changed + } + l.used++ + return &carrierLease{limiter: l}, nil +} + +func (l *CarrierLimiter) release() { + l.mu.Lock() + l.used-- + changed := l.changed + l.changed = make(chan struct{}) + l.mu.Unlock() + close(changed) +} + +type carrierLease struct { + limiter *CarrierLimiter + once sync.Once +} + +func (l *carrierLease) release() { + if l == nil || l.limiter == nil { + return + } + l.once.Do(l.limiter.release) +} + +type limitedCarrier struct { + net.Conn + lease *carrierLease +} + +func (c *limitedCarrier) Close() error { + err := c.Conn.Close() + c.lease.release() + return err +} diff --git a/transport/muxcool/codec.go b/transport/muxcool/codec.go new file mode 100644 index 0000000000..64aa724ac0 --- /dev/null +++ b/transport/muxcool/codec.go @@ -0,0 +1,454 @@ +package muxcool + +import ( + "encoding/binary" + "errors" + "fmt" + "io" + "net" + "net/netip" + "strings" + + "github.com/metacubex/mihomo/common/pool" +) + +const ( + MaxMetadataSize = 512 + MaxPayloadSize = 8 * 1024 +) + +type Status byte + +const ( + StatusNew Status = 0x01 + StatusKeep Status = 0x02 + StatusEnd Status = 0x03 + StatusKeepAlive Status = 0x04 +) + +type Option byte + +const ( + OptionData Option = 0x01 + OptionError Option = 0x02 +) + +type Network byte + +const ( + NetworkTCP Network = 0x01 + NetworkUDP Network = 0x02 +) + +const ( + addressIPv4 byte = 0x01 + addressDomain byte = 0x02 + addressIPv6 byte = 0x03 +) + +type Frame struct { + SessionID uint16 + Status Status + Option Option + Network Network + Destination string + DestinationIP netip.Addr + Port uint16 + GlobalID [8]byte + Payload []byte +} + +type decodedFrame struct { + Frame + payloadPooled bool +} + +type ProtocolError struct { + Op string + Err error +} + +func (e *ProtocolError) Error() string { + return fmt.Sprintf("mux.cool %s: %v", e.Op, e.Err) +} + +func (e *ProtocolError) Unwrap() error { + return e.Err +} + +func protocolError(op string, err error) error { + return &ProtocolError{Op: op, Err: err} +} + +func EncodeFrame(frame Frame) ([]byte, error) { + return encodeFrame(nil, frame) +} + +func encodeFrame(buffer []byte, frame Frame) ([]byte, error) { + if frame.Status == StatusKeep && + frame.Option == OptionData && + frame.Network == 0 && + frame.Destination == "" && + !frame.DestinationIP.IsValid() && + frame.Port == 0 && + frame.GlobalID == [8]byte{} && + len(frame.Payload) <= int(^uint16(0)) { + return encodeKeepDataFrame(buffer, frame.SessionID, frame.Payload), nil + } + if err := validateFrame(frame); err != nil { + return nil, protocolError("encode", err) + } + + hasTarget := frame.Status == StatusNew || (frame.Status == StatusKeep && (frame.Destination != "" || frame.DestinationIP.IsValid())) + metadataLen := 4 + targetAddr := frame.DestinationIP.Unmap() + if hasTarget { + metadataLen += 3 + if targetAddr.IsValid() && targetAddr.Zone() != "" { + return nil, protocolError("encode address", errors.New("scoped IPv6 addresses are not supported")) + } + if parsed, ok := parseLiteralIP(frame.Destination); !targetAddr.IsValid() && ok { + targetAddr = parsed.Unmap() + } + if targetAddr.IsValid() { + if targetAddr.Is4() { + metadataLen += 1 + net.IPv4len + } else { + metadataLen += 1 + net.IPv6len + } + } else { + if len(frame.Destination) == 0 || len(frame.Destination) > 255 { + return nil, protocolError("encode address", fmt.Errorf("invalid domain length %d", len(frame.Destination))) + } + metadataLen += 2 + len(frame.Destination) + } + } + if frame.Status == StatusNew && frame.Network == NetworkUDP && frame.Option&OptionData != 0 { + metadataLen += len(frame.GlobalID) + } + if metadataLen > MaxMetadataSize { + return nil, protocolError("encode", fmt.Errorf("metadata length %d exceeds %d", metadataLen, MaxMetadataSize)) + } + + frameLen := 2 + metadataLen + if frame.Option&OptionData != 0 { + frameLen += 2 + len(frame.Payload) + } + var result []byte + if cap(buffer) >= frameLen { + result = buffer[:frameLen] + } else { + result = make([]byte, frameLen) + } + binary.BigEndian.PutUint16(result, uint16(metadataLen)) + offset := 2 + binary.BigEndian.PutUint16(result[offset:], frame.SessionID) + offset += 2 + result[offset] = byte(frame.Status) + offset++ + result[offset] = byte(frame.Option) + offset++ + if hasTarget { + result[offset] = byte(frame.Network) + offset++ + binary.BigEndian.PutUint16(result[offset:], frame.Port) + offset += 2 + if targetAddr.IsValid() { + if targetAddr.Is4() { + result[offset] = addressIPv4 + offset++ + address := targetAddr.As4() + offset += copy(result[offset:], address[:]) + } else { + result[offset] = addressIPv6 + offset++ + address := targetAddr.As16() + offset += copy(result[offset:], address[:]) + } + } else { + result[offset] = addressDomain + result[offset+1] = byte(len(frame.Destination)) + offset += 2 + offset += copy(result[offset:], frame.Destination) + } + } + if frame.Status == StatusNew && frame.Network == NetworkUDP && frame.Option&OptionData != 0 { + offset += copy(result[offset:], frame.GlobalID[:]) + } + if frame.Option&OptionData != 0 { + binary.BigEndian.PutUint16(result[offset:], uint16(len(frame.Payload))) + offset += 2 + copy(result[offset:], frame.Payload) + } + return result, nil +} + +func encodeKeepDataFrame(buffer []byte, sessionID uint16, payload []byte) []byte { + frameLen := 8 + len(payload) + var result []byte + if cap(buffer) >= frameLen { + result = buffer[:frameLen] + } else { + result = make([]byte, frameLen) + } + binary.BigEndian.PutUint16(result, 4) + binary.BigEndian.PutUint16(result[2:], sessionID) + result[4] = byte(StatusKeep) + result[5] = byte(OptionData) + binary.BigEndian.PutUint16(result[6:], uint16(len(payload))) + copy(result[8:], payload) + return result +} + +func parseLiteralIP(host string) (netip.Addr, bool) { + if host == "" { + return netip.Addr{}, false + } + first := host[0] + if (first < '0' || first > '9') && strings.IndexByte(host, ':') < 0 { + return netip.Addr{}, false + } + addr, err := netip.ParseAddr(host) + if err != nil || addr.Zone() != "" { + return netip.Addr{}, false + } + return addr.Unmap(), true +} + +func DecodeFrame(r io.Reader) (Frame, error) { + return decodeFrame(r, nil) +} + +func decodeFrame(r io.Reader, metadataBuffer []byte) (Frame, error) { + decoded, err := decodeFrameWithPayloadPool(r, metadataBuffer, false) + return decoded.Frame, err +} + +func decodeFramePooled(r io.Reader, metadataBuffer []byte) (decodedFrame, error) { + return decodeFrameWithPayloadPool(r, metadataBuffer, true) +} + +func decodeFrameWithPayloadPool(r io.Reader, metadataBuffer []byte, poolPayload bool) (decodedFrame, error) { + var lengthBytes []byte + if len(metadataBuffer) >= 2 { + lengthBytes = metadataBuffer[:2] + } else { + lengthBytes = make([]byte, 2) + } + if _, err := io.ReadFull(r, lengthBytes); err != nil { + return decodedFrame{}, protocolError("read metadata length", err) + } + metadataLen := int(binary.BigEndian.Uint16(lengthBytes)) + if metadataLen < 4 || metadataLen > MaxMetadataSize { + return decodedFrame{}, protocolError("read metadata", fmt.Errorf("invalid metadata length %d", metadataLen)) + } + var metadata []byte + if len(metadataBuffer) >= metadataLen { + metadata = metadataBuffer[:metadataLen] + } else { + metadata = make([]byte, metadataLen) + } + if _, err := io.ReadFull(r, metadata); err != nil { + return decodedFrame{}, protocolError("read metadata", err) + } + + frame := Frame{ + SessionID: binary.BigEndian.Uint16(metadata[:2]), + Status: Status(metadata[2]), + Option: Option(metadata[3]), + } + fastKeepData := metadataLen == 4 && frame.Status == StatusKeep && frame.Option == OptionData + if !fastKeepData { + if frame.Status != StatusNew && frame.Status != StatusKeep && frame.Status != StatusEnd && frame.Status != StatusKeepAlive { + return decodedFrame{}, protocolError("decode metadata", fmt.Errorf("invalid status %d", frame.Status)) + } + if frame.Option & ^(OptionData|OptionError) != 0 { + return decodedFrame{}, protocolError("decode metadata", fmt.Errorf("invalid option %d", frame.Option)) + } + + targetBytes := metadata[4:] + hasTarget := frame.Status == StatusNew || (frame.Status == StatusKeep && len(targetBytes) > 0) + if len(targetBytes) > 0 && !hasTarget { + return decodedFrame{}, protocolError("decode metadata", fmt.Errorf("unexpected trailing metadata: %d bytes", len(targetBytes))) + } + if hasTarget { + if len(targetBytes) < 4 { + return decodedFrame{}, protocolError("decode target", io.ErrUnexpectedEOF) + } + frame.Network = Network(targetBytes[0]) + if frame.Network != NetworkTCP && frame.Network != NetworkUDP { + return decodedFrame{}, protocolError("decode target", fmt.Errorf("invalid network %d", frame.Network)) + } + if frame.Status == StatusKeep && frame.Network != NetworkUDP { + return decodedFrame{}, protocolError("decode target", errors.New("follow-up target is only valid for UDP")) + } + frame.Port = binary.BigEndian.Uint16(targetBytes[1:3]) + host, address, consumed, err := readAddress(targetBytes[3:], poolPayload) + if err != nil { + return decodedFrame{}, protocolError("decode target", err) + } + frame.Destination = host + frame.DestinationIP = address + targetBytes = targetBytes[3+consumed:] + if frame.Status == StatusNew && frame.Network == NetworkUDP && frame.Option&OptionData != 0 { + if len(targetBytes) != len(frame.GlobalID) { + return decodedFrame{}, protocolError("decode GlobalID", fmt.Errorf("invalid length %d", len(targetBytes))) + } + copy(frame.GlobalID[:], targetBytes) + targetBytes = nil + } + if len(targetBytes) != 0 { + return decodedFrame{}, protocolError("decode target", fmt.Errorf("unexpected trailing metadata: %d bytes", len(targetBytes))) + } + } + } + + if frame.Option&OptionData != 0 { + if _, err := io.ReadFull(r, lengthBytes); err != nil { + return decodedFrame{}, protocolError("read payload length", err) + } + payloadLen := int(binary.BigEndian.Uint16(lengthBytes)) + payloadPooled := false + if poolPayload && payloadLen > 0 { + frame.Payload = pool.Get(payloadLen)[:payloadLen] + payloadPooled = true + } else { + frame.Payload = make([]byte, payloadLen) + } + if _, err := io.ReadFull(r, frame.Payload); err != nil { + decoded := decodedFrame{Frame: frame, payloadPooled: payloadPooled} + decoded.releasePayload() + return decodedFrame{}, protocolError("read payload", err) + } + return decodedFrame{Frame: frame, payloadPooled: payloadPooled}, nil + } + return decodedFrame{Frame: frame}, nil +} + +func (f *decodedFrame) releasePayload() { + releasePooledPayload(f.Payload, f.payloadPooled) + f.Payload = nil + f.payloadPooled = false +} + +func releasePooledPayload(payload []byte, pooled bool) { + if pooled { + _ = pool.Put(payload) + } +} + +func validateFrame(frame Frame) error { + switch frame.Status { + case StatusNew, StatusKeep, StatusEnd, StatusKeepAlive: + default: + return fmt.Errorf("invalid status %d", frame.Status) + } + if frame.Option & ^(OptionData|OptionError) != 0 { + return fmt.Errorf("invalid option %d", frame.Option) + } + if len(frame.Payload) > int(^uint16(0)) { + return fmt.Errorf("payload length %d exceeds uint16", len(frame.Payload)) + } + if len(frame.Payload) > 0 && frame.Option&OptionData == 0 { + return errors.New("payload provided without data option") + } + if frame.GlobalID != [8]byte{} && (frame.Status != StatusNew || frame.Network != NetworkUDP || frame.Option&OptionData == 0) { + return errors.New("GlobalID is only valid on an initial UDP data frame") + } + hasTarget := frame.Destination != "" || frame.DestinationIP.IsValid() + if frame.Status == StatusNew || (frame.Status == StatusKeep && hasTarget) { + if frame.Network != NetworkTCP && frame.Network != NetworkUDP { + return fmt.Errorf("invalid network %d", frame.Network) + } + if !hasTarget { + return errors.New("empty destination") + } + if frame.Status == StatusKeep && frame.Network != NetworkUDP { + return errors.New("follow-up target is only valid for UDP") + } + } + return nil +} + +func readAddress(raw []byte, structuredIP bool) (string, netip.Addr, int, error) { + if len(raw) < 1 { + return "", netip.Addr{}, 0, io.ErrUnexpectedEOF + } + switch raw[0] { + case addressIPv4: + if len(raw) < 1+net.IPv4len { + return "", netip.Addr{}, 0, io.ErrUnexpectedEOF + } + var rawAddress [net.IPv4len]byte + copy(rawAddress[:], raw[1:1+net.IPv4len]) + address := netip.AddrFrom4(rawAddress) + if structuredIP { + return "", address, 1 + net.IPv4len, nil + } + return address.String(), netip.Addr{}, 1 + net.IPv4len, nil + case addressIPv6: + if len(raw) < 1+net.IPv6len { + return "", netip.Addr{}, 0, io.ErrUnexpectedEOF + } + var rawAddress [net.IPv6len]byte + copy(rawAddress[:], raw[1:1+net.IPv6len]) + address := netip.AddrFrom16(rawAddress) + if structuredIP { + return "", address, 1 + net.IPv6len, nil + } + return address.String(), netip.Addr{}, 1 + net.IPv6len, nil + case addressDomain: + if len(raw) < 2 { + return "", netip.Addr{}, 0, io.ErrUnexpectedEOF + } + length := int(raw[1]) + if length == 0 || len(raw) < 2+length { + return "", netip.Addr{}, 0, io.ErrUnexpectedEOF + } + return string(raw[2 : 2+length]), netip.Addr{}, 2 + length, nil + default: + return "", netip.Addr{}, 0, fmt.Errorf("invalid address type %d", raw[0]) + } +} + +func writeStreamData(w io.Writer, sessionID uint16, destination string, port uint16, initial bool, payload []byte) error { + first := initial + for len(payload) > 0 { + chunkSize := len(payload) + if chunkSize > MaxPayloadSize { + chunkSize = MaxPayloadSize + } + status := StatusKeep + frame := Frame{SessionID: sessionID, Status: status, Option: OptionData, Payload: payload[:chunkSize]} + if first { + frame.Status = StatusNew + frame.Network = NetworkTCP + frame.Destination = destination + frame.Port = port + first = false + } + encoded, err := EncodeFrame(frame) + if err != nil { + return err + } + if err := writeFull(w, encoded); err != nil { + return err + } + payload = payload[chunkSize:] + } + return nil +} + +func writeFull(w io.Writer, payload []byte) error { + for len(payload) > 0 { + n, err := w.Write(payload) + if err != nil { + return err + } + if n == 0 { + return io.ErrShortWrite + } + payload = payload[n:] + } + return nil +} diff --git a/transport/muxcool/codec_decode_test.go b/transport/muxcool/codec_decode_test.go new file mode 100644 index 0000000000..f4431ff47d --- /dev/null +++ b/transport/muxcool/codec_decode_test.go @@ -0,0 +1,122 @@ +package muxcool + +import ( + "bytes" + "encoding/hex" + "errors" + "net/netip" + "testing" +) + +func TestDecodeFrameReadsXrayGoldenVectors(t *testing.T) { + tests := []struct { + name string + status Status + option Option + sessionID uint16 + network Network + destination string + port uint16 + payload string + }{ + {name: "domain-new-data", status: StatusNew, option: OptionData, sessionID: 1, network: NetworkTCP, destination: "example.com", port: 443, payload: "hi"}, + {name: "domain-new-empty", status: StatusNew, sessionID: 1, network: NetworkTCP, destination: "example.com", port: 443}, + {name: "ipv4-new-empty", status: StatusNew, sessionID: 2, network: NetworkTCP, destination: "1.2.3.4", port: 53}, + {name: "ipv6-new-empty", status: StatusNew, sessionID: 3, network: NetworkTCP, destination: "::1", port: 8080}, + {name: "keep-data", status: StatusKeep, option: OptionData, sessionID: 1, payload: "ok"}, + {name: "udp-new-data", status: StatusNew, option: OptionData, sessionID: 4, network: NetworkUDP, destination: "udp.example", port: 53, payload: "q1"}, + {name: "xudp-new-data", status: StatusNew, option: OptionData, sessionID: 5, network: NetworkUDP, destination: "udp.example", port: 53, payload: "q2"}, + {name: "udp-keep-data", status: StatusKeep, option: OptionData, sessionID: 4, network: NetworkUDP, destination: "udp.example", port: 53, payload: "r3"}, + {name: "end", status: StatusEnd, sessionID: 1}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + raw, err := hex.DecodeString(referenceFrames[tt.name]) + if err != nil { + t.Fatalf("decode fixture: %v", err) + } + got, err := DecodeFrame(bytes.NewReader(raw)) + if err != nil { + t.Fatalf("DecodeFrame: %v", err) + } + if got.Status != tt.status || got.Option != tt.option || got.SessionID != tt.sessionID { + t.Fatalf("header = status %d option %d session %d", got.Status, got.Option, got.SessionID) + } + if got.Network != tt.network || got.Destination != tt.destination || got.Port != tt.port { + t.Fatalf("target = network %d %s:%d", got.Network, got.Destination, got.Port) + } + if string(got.Payload) != tt.payload { + t.Fatalf("payload = %q, want %q", got.Payload, tt.payload) + } + if tt.name == "xudp-new-data" && got.GlobalID != [8]byte{1, 2, 3, 4, 5, 6, 7, 8} { + t.Fatalf("GlobalID = %v", got.GlobalID) + } + }) + } +} + +func TestDecodeFramePooledKeepsStructuredIP(t *testing.T) { + raw, err := EncodeFrame(Frame{ + SessionID: 1, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, + DestinationIP: netip.MustParseAddr("1.2.3.4"), Port: 53, Payload: []byte("payload"), + }) + if err != nil { + t.Fatal(err) + } + + decoded, err := decodeFramePooled(bytes.NewReader(raw), make([]byte, MaxMetadataSize)) + if err != nil { + t.Fatal(err) + } + defer decoded.releasePayload() + if decoded.Destination != "" || decoded.DestinationIP != netip.MustParseAddr("1.2.3.4") { + t.Fatalf("decoded target = (%q, %v), want structured IPv4", decoded.Destination, decoded.DestinationIP) + } + + public, err := DecodeFrame(bytes.NewReader(raw)) + if err != nil { + t.Fatal(err) + } + if public.Destination != "1.2.3.4" || public.DestinationIP.IsValid() { + t.Fatalf("public target = (%q, %v), want stable string representation", public.Destination, public.DestinationIP) + } +} + +func TestDecodeFrameRejectsMalformedInput(t *testing.T) { + tests := []struct { + name string + raw []byte + }{ + {name: "truncated length", raw: []byte{0}}, + {name: "oversized metadata", raw: []byte{0x02, 0x01}}, + {name: "metadata shorter than header", raw: []byte{0, 3, 0, 1, 1}}, + {name: "unknown status", raw: []byte{0, 4, 0, 1, 9, 0}}, + {name: "unknown option", raw: []byte{0, 4, 0, 1, 2, 0x80}}, + {name: "new missing target", raw: []byte{0, 4, 0, 1, 1, 0}}, + {name: "invalid network", raw: []byte{0, 8, 0, 1, 1, 0, 9, 0, 80, 2}}, + {name: "invalid address type", raw: []byte{0, 8, 0, 1, 1, 0, 1, 0, 80, 9}}, + {name: "truncated domain", raw: []byte{0, 10, 0, 1, 1, 0, 1, 0, 80, 2, 4, 'a'}}, + {name: "missing payload length", raw: []byte{0, 4, 0, 1, 2, 1}}, + {name: "truncated payload", raw: []byte{0, 4, 0, 1, 2, 1, 0, 4, 'a'}}, + {name: "new UDP data missing GlobalID", raw: []byte{0, 12, 0, 1, 1, 1, 2, 0, 53, 1, 1, 2, 3, 4, 0, 1, 'x'}}, + {name: "new UDP data partial GlobalID", raw: []byte{0, 16, 0, 1, 1, 1, 2, 0, 53, 1, 1, 2, 3, 4, 1, 2, 3, 4, 0, 1, 'x'}}, + {name: "keep UDP trailing metadata", raw: []byte{0, 13, 0, 1, 2, 1, 2, 0, 53, 1, 1, 2, 3, 4, 0, 0xff, 0, 1, 'x'}}, + {name: "keep TCP target", raw: []byte{0, 12, 0, 1, 2, 1, 1, 0, 53, 1, 1, 2, 3, 4, 0, 1, 'x'}}, + {name: "end with target", raw: []byte{0, 12, 0, 1, 3, 0, 2, 0, 53, 1, 1, 2, 3, 4}}, + {name: "keepalive with metadata", raw: []byte{0, 5, 0, 1, 4, 0, 0}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := DecodeFrame(bytes.NewReader(tt.raw)) + if err == nil { + t.Fatal("expected protocol error") + } + var protocolErr *ProtocolError + if !errors.As(err, &protocolErr) { + t.Fatalf("error type = %T, want *ProtocolError: %v", err, err) + } + }) + } +} diff --git a/transport/muxcool/codec_encode_test.go b/transport/muxcool/codec_encode_test.go new file mode 100644 index 0000000000..050ba2203f --- /dev/null +++ b/transport/muxcool/codec_encode_test.go @@ -0,0 +1,137 @@ +package muxcool + +import ( + "bytes" + "encoding/hex" + "testing" +) + +func TestEncodeFrameMatchesXrayGoldenVectors(t *testing.T) { + tests := []struct { + name string + frame Frame + }{ + { + name: "domain-new-data", + frame: Frame{SessionID: 1, Status: StatusNew, Option: OptionData, Network: NetworkTCP, + Destination: "example.com", Port: 443, Payload: []byte("hi")}, + }, + { + name: "domain-new-empty", + frame: Frame{SessionID: 1, Status: StatusNew, Network: NetworkTCP, + Destination: "example.com", Port: 443}, + }, + { + name: "ipv4-new-empty", + frame: Frame{SessionID: 2, Status: StatusNew, Network: NetworkTCP, + Destination: "1.2.3.4", Port: 53}, + }, + { + name: "ipv6-new-empty", + frame: Frame{SessionID: 3, Status: StatusNew, Network: NetworkTCP, + Destination: "::1", Port: 8080}, + }, + { + name: "keep-data", + frame: Frame{SessionID: 1, Status: StatusKeep, Option: OptionData, Payload: []byte("ok")}, + }, + { + name: "udp-new-data", + frame: Frame{SessionID: 4, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "udp.example", Port: 53, Payload: []byte("q1")}, + }, + { + name: "xudp-new-data", + frame: Frame{SessionID: 5, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "udp.example", Port: 53, GlobalID: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}, Payload: []byte("q2")}, + }, + { + name: "udp-keep-data", + frame: Frame{SessionID: 4, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, + Destination: "udp.example", Port: 53, Payload: []byte("r3")}, + }, + { + name: "end", + frame: Frame{SessionID: 1, Status: StatusEnd}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := EncodeFrame(tt.frame) + if err != nil { + t.Fatalf("EncodeFrame: %v", err) + } + want, err := hex.DecodeString(referenceFrames[tt.name]) + if err != nil { + t.Fatalf("decode fixture: %v", err) + } + if !bytes.Equal(got, want) { + t.Fatalf("frame mismatch\n got: %x\nwant: %x", got, want) + } + }) + } +} + +func TestEncodeFrameRejectsGlobalIDOutsideInitialUDPPayload(t *testing.T) { + globalID := [8]byte{1} + for _, frame := range []Frame{ + {SessionID: 1, Status: StatusNew, Network: NetworkTCP, Destination: "tcp.example", Port: 443, GlobalID: globalID}, + {SessionID: 1, Status: StatusNew, Network: NetworkUDP, Destination: "udp.example", Port: 53, GlobalID: globalID}, + {SessionID: 1, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, Destination: "udp.example", Port: 53, GlobalID: globalID, Payload: []byte("x")}, + } { + if _, err := EncodeFrame(frame); err == nil { + t.Fatalf("EncodeFrame(%+v) accepted invalid GlobalID", frame) + } + } +} + +func TestWriteStreamDataChunksAtEightKiB(t *testing.T) { + payload := bytes.Repeat([]byte{0x5a}, MaxPayloadSize+17) + var carrier bytes.Buffer + + if err := writeStreamData(&carrier, 7, "large.example", 8443, true, payload); err != nil { + t.Fatalf("writeStreamData: %v", err) + } + + first, err := decodeReferenceFrameFrom(&carrier) + if err != nil { + t.Fatalf("decode first frame: %v", err) + } + second, err := decodeReferenceFrameFrom(&carrier) + if err != nil { + t.Fatalf("decode second frame: %v", err) + } + if first.status != byte(StatusNew) || len(first.payload) != MaxPayloadSize { + t.Fatalf("first frame status=%d payload=%d", first.status, len(first.payload)) + } + if second.status != byte(StatusKeep) || len(second.payload) != 17 { + t.Fatalf("second frame status=%d payload=%d", second.status, len(second.payload)) + } + if carrier.Len() != 0 { + t.Fatalf("unexpected trailing bytes: %d", carrier.Len()) + } +} + +func decodeReferenceFrameFrom(carrier *bytes.Buffer) (referenceFrame, error) { + if carrier.Len() < 2 { + return referenceFrame{}, bytes.ErrTooLarge + } + metaLen := int(carrier.Bytes()[0])<<8 | int(carrier.Bytes()[1]) + total := 2 + metaLen + if carrier.Len() < total { + return referenceFrame{}, bytes.ErrTooLarge + } + if carrier.Bytes()[5]&byte(OptionData) != 0 { + if carrier.Len() < total+2 { + return referenceFrame{}, bytes.ErrTooLarge + } + payloadLen := int(carrier.Bytes()[total])<<8 | int(carrier.Bytes()[total+1]) + total += 2 + payloadLen + } + if carrier.Len() < total { + return referenceFrame{}, bytes.ErrTooLarge + } + raw := append([]byte(nil), carrier.Next(total)...) + return decodeReferenceFrame(raw) +} diff --git a/transport/muxcool/fixtures_test.go b/transport/muxcool/fixtures_test.go new file mode 100644 index 0000000000..1b42142018 --- /dev/null +++ b/transport/muxcool/fixtures_test.go @@ -0,0 +1,192 @@ +package muxcool + +import ( + "bytes" + "encoding/binary" + "encoding/hex" + "errors" + "io" + "net" + "sync" + "testing" + "time" +) + +// Golden vectors are derived independently from Xray-core common/mux at +// revision 0ee156e75c9546a713f6c88c0bd14f5ff953c567. Keep these literals +// independent from the production encoder so they detect wire-format drift. +var referenceFrames = map[string]string{ + "domain-new-data": "0014000101010101bb020b6578616d706c652e636f6d00026869", + "domain-new-empty": "0014000101000101bb020b6578616d706c652e636f6d", + "ipv4-new-empty": "000c000201000100350101020304", + "ipv6-new-empty": "001800030100011f900300000000000000000000000000000001", + "keep-data": "00040001020100026f6b", + "udp-new-data": "001c00040101020035020b7564702e6578616d706c65000000000000000000027131", + "xudp-new-data": "001c00050101020035020b7564702e6578616d706c65010203040506070800027132", + "udp-keep-data": "001400040201020035020b7564702e6578616d706c6500027233", + "end": "000400010300", +} + +type referenceFrame struct { + metaLen uint16 + sessionID uint16 + status byte + option byte + metadata []byte + payload []byte +} + +func decodeReferenceFrame(raw []byte) (referenceFrame, error) { + if len(raw) < 6 { + return referenceFrame{}, io.ErrUnexpectedEOF + } + metaLen := binary.BigEndian.Uint16(raw[:2]) + if metaLen < 4 || int(metaLen)+2 > len(raw) { + return referenceFrame{}, io.ErrUnexpectedEOF + } + meta := raw[2 : 2+metaLen] + frame := referenceFrame{ + metaLen: metaLen, + sessionID: binary.BigEndian.Uint16(meta[:2]), + status: meta[2], + option: meta[3], + metadata: append([]byte(nil), meta[4:]...), + } + if frame.option&1 == 0 { + if len(raw) != int(metaLen)+2 { + return referenceFrame{}, errors.New("unexpected bytes after metadata-only frame") + } + return frame, nil + } + offset := 2 + int(metaLen) + if len(raw) < offset+2 { + return referenceFrame{}, io.ErrUnexpectedEOF + } + payloadLen := int(binary.BigEndian.Uint16(raw[offset : offset+2])) + if len(raw) != offset+2+payloadLen { + return referenceFrame{}, io.ErrUnexpectedEOF + } + frame.payload = append([]byte(nil), raw[offset+2:]...) + return frame, nil +} + +func TestReferenceFrameFixtures(t *testing.T) { + tests := []struct { + name string + status byte + option byte + sessionID uint16 + payload string + }{ + {name: "domain-new-data", status: 1, option: 1, sessionID: 1, payload: "hi"}, + {name: "domain-new-empty", status: 1, sessionID: 1}, + {name: "ipv4-new-empty", status: 1, sessionID: 2}, + {name: "ipv6-new-empty", status: 1, sessionID: 3}, + {name: "keep-data", status: 2, option: 1, sessionID: 1, payload: "ok"}, + {name: "udp-new-data", status: 1, option: 1, sessionID: 4, payload: "q1"}, + {name: "xudp-new-data", status: 1, option: 1, sessionID: 5, payload: "q2"}, + {name: "udp-keep-data", status: 2, option: 1, sessionID: 4, payload: "r3"}, + {name: "end", status: 3, sessionID: 1}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + raw, err := hex.DecodeString(referenceFrames[tt.name]) + if err != nil { + t.Fatalf("decode fixture: %v", err) + } + frame, err := decodeReferenceFrame(raw) + if err != nil { + t.Fatalf("decode reference frame: %v", err) + } + if frame.status != tt.status || frame.option != tt.option || frame.sessionID != tt.sessionID { + t.Fatalf("header = status %d option %d session %d", frame.status, frame.option, frame.sessionID) + } + if string(frame.payload) != tt.payload { + t.Fatalf("payload = %q, want %q", frame.payload, tt.payload) + } + }) + } +} + +type fakeTimer struct { + mu sync.Mutex + stopped bool + fn func() +} + +func (t *fakeTimer) Stop() bool { + t.mu.Lock() + defer t.mu.Unlock() + wasActive := !t.stopped + t.stopped = true + return wasActive +} + +func (t *fakeTimer) fire() { + t.mu.Lock() + if t.stopped { + t.mu.Unlock() + return + } + t.stopped = true + fn := t.fn + t.mu.Unlock() + fn() +} + +type fakeClock struct { + mu sync.Mutex + timers []*fakeTimer +} + +func (c *fakeClock) AfterFunc(_ time.Duration, fn func()) *fakeTimer { + t := &fakeTimer{fn: fn} + c.mu.Lock() + c.timers = append(c.timers, t) + c.mu.Unlock() + return t +} + +func (c *fakeClock) FireAll() { + c.mu.Lock() + timers := append([]*fakeTimer(nil), c.timers...) + c.timers = nil + c.mu.Unlock() + for _, timer := range timers { + timer.fire() + } +} + +func (c *fakeClock) timerCount() int { + c.mu.Lock() + defer c.mu.Unlock() + return len(c.timers) +} + +type recordingConn struct { + net.Conn + mu sync.Mutex + writes bytes.Buffer + closed bool +} + +func (c *recordingConn) Write(p []byte) (int, error) { + c.mu.Lock() + _, _ = c.writes.Write(p) + c.mu.Unlock() + return c.Conn.Write(p) +} + +func (c *recordingConn) Close() error { + c.mu.Lock() + c.closed = true + c.mu.Unlock() + return c.Conn.Close() +} + +func (c *recordingConn) Bytes() []byte { + c.mu.Lock() + defer c.mu.Unlock() + return append([]byte(nil), c.writes.Bytes()...) +} diff --git a/transport/muxcool/integration_test.go b/transport/muxcool/integration_test.go new file mode 100644 index 0000000000..2ba7281460 --- /dev/null +++ b/transport/muxcool/integration_test.go @@ -0,0 +1,391 @@ +package muxcool + +import ( + "context" + "encoding/binary" + "fmt" + "io" + "net" + "sync" + "sync/atomic" + "testing" + "time" +) + +type independentFrame struct { + sessionID uint16 + status byte + option byte + network byte + host string + port uint16 + globalID [8]byte + payload []byte +} + +func readIndependentFrame(reader io.Reader) (independentFrame, error) { + var lengthBytes [2]byte + if _, err := io.ReadFull(reader, lengthBytes[:]); err != nil { + return independentFrame{}, err + } + metadataLength := int(binary.BigEndian.Uint16(lengthBytes[:])) + if metadataLength < 4 || metadataLength > 512 { + return independentFrame{}, fmt.Errorf("invalid metadata length %d", metadataLength) + } + metadata := make([]byte, metadataLength) + if _, err := io.ReadFull(reader, metadata); err != nil { + return independentFrame{}, err + } + frame := independentFrame{ + sessionID: binary.BigEndian.Uint16(metadata[:2]), + status: metadata[2], + option: metadata[3], + } + if frame.status == 1 || (frame.status == 2 && len(metadata) > 4) { + if len(metadata) < 8 { + return independentFrame{}, io.ErrUnexpectedEOF + } + frame.network = metadata[4] + frame.port = binary.BigEndian.Uint16(metadata[5:7]) + host, consumed, err := readIndependentAddress(metadata[7:]) + if err != nil { + return independentFrame{}, err + } + frame.host = host + remaining := metadata[7+consumed:] + if frame.status == 1 && frame.network == 2 && frame.option&1 != 0 { + if len(remaining) != len(frame.globalID) { + return independentFrame{}, fmt.Errorf("invalid GlobalID length %d", len(remaining)) + } + copy(frame.globalID[:], remaining) + } else if len(remaining) != 0 { + return independentFrame{}, fmt.Errorf("unexpected trailing metadata %d", len(remaining)) + } + } + if frame.option&1 != 0 { + if _, err := io.ReadFull(reader, lengthBytes[:]); err != nil { + return independentFrame{}, err + } + frame.payload = make([]byte, int(binary.BigEndian.Uint16(lengthBytes[:]))) + if _, err := io.ReadFull(reader, frame.payload); err != nil { + return independentFrame{}, err + } + } + return frame, nil +} + +func readIndependentAddress(raw []byte) (string, int, error) { + if len(raw) == 0 { + return "", 0, io.ErrUnexpectedEOF + } + switch raw[0] { + case 1: + if len(raw) < 5 { + return "", 0, io.ErrUnexpectedEOF + } + return net.IP(raw[1:5]).String(), 5, nil + case 2: + if len(raw) < 2 || len(raw) < 2+int(raw[1]) { + return "", 0, io.ErrUnexpectedEOF + } + return string(raw[2 : 2+int(raw[1])]), 2 + int(raw[1]), nil + case 3: + if len(raw) < 17 { + return "", 0, io.ErrUnexpectedEOF + } + return net.IP(raw[1:17]).String(), 17, nil + default: + return "", 0, fmt.Errorf("invalid address type %d", raw[0]) + } +} + +func writeIndependentKeep(writer io.Writer, sessionID uint16, payload []byte) error { + raw := make([]byte, 0, 8+len(payload)) + raw = binary.BigEndian.AppendUint16(raw, 4) + raw = binary.BigEndian.AppendUint16(raw, sessionID) + raw = append(raw, 2, 1) + raw = binary.BigEndian.AppendUint16(raw, uint16(len(payload))) + raw = append(raw, payload...) + _, err := writer.Write(raw) + return err +} + +func writeIndependentUDPKeep(writer io.Writer, frame independentFrame) error { + metadata := make([]byte, 0, 32) + metadata = binary.BigEndian.AppendUint16(metadata, frame.sessionID) + metadata = append(metadata, 2, 1, 2) + metadata = binary.BigEndian.AppendUint16(metadata, frame.port) + if ip := net.ParseIP(frame.host); ip != nil { + if ip4 := ip.To4(); ip4 != nil { + metadata = append(metadata, 1) + metadata = append(metadata, ip4...) + } else { + metadata = append(metadata, 3) + metadata = append(metadata, ip.To16()...) + } + } else { + metadata = append(metadata, 2, byte(len(frame.host))) + metadata = append(metadata, frame.host...) + } + raw := make([]byte, 0, 4+len(metadata)+len(frame.payload)) + raw = binary.BigEndian.AppendUint16(raw, uint16(len(metadata))) + raw = append(raw, metadata...) + raw = binary.BigEndian.AppendUint16(raw, uint16(len(frame.payload))) + raw = append(raw, frame.payload...) + _, err := writer.Write(raw) + return err +} + +func serveIndependentEcho(carrier net.Conn) { + defer carrier.Close() + for { + frame, err := readIndependentFrame(carrier) + if err != nil { + return + } + if frame.status == 3 || frame.option&1 == 0 { + continue + } + if err := writeIndependentKeep(carrier, frame.sessionID, frame.payload); err != nil { + return + } + } +} + +func serveIndependentPacketEcho(carrier net.Conn, observed chan<- independentFrame) { + defer carrier.Close() + for { + frame, err := readIndependentFrame(carrier) + if err != nil { + return + } + if frame.status == 3 || frame.option&1 == 0 { + continue + } + if observed != nil { + observed <- frame + } + if err := writeIndependentUDPKeep(carrier, frame); err != nil { + return + } + } +} + +func TestPoolInteroperatesWithIndependentXrayCompatibleEchoServer(t *testing.T) { + var carrierDials atomic.Int32 + pool := NewPool(func(context.Context) (net.Conn, error) { + carrierDials.Add(1) + client, server := net.Pipe() + go serveIndependentEcho(server) + return client, nil + }, Options{ + MaxConcurrency: 16, + MaxConnections: 128, + FirstPayloadTimeout: time.Hour, + IdleTimeout: time.Hour, + }) + t.Cleanup(func() { _ = pool.Close() }) + + const streamCount = 12 + connections := make([]net.Conn, 0, streamCount) + for index := 0; index < streamCount; index++ { + conn, err := pool.DialContext(context.Background(), fmt.Sprintf("echo-%d.example", index), uint16(8000+index)) + if err != nil { + t.Fatalf("dial stream %d: %v", index, err) + } + connections = append(connections, conn) + } + + var wg sync.WaitGroup + for index, conn := range connections { + wg.Add(1) + go func(index int, conn net.Conn) { + defer wg.Done() + defer conn.Close() + request := []byte(fmt.Sprintf("stream-%d", index)) + if _, err := conn.Write(request); err != nil { + t.Errorf("write stream %d: %v", index, err) + return + } + response := make([]byte, len(request)) + if _, err := io.ReadFull(conn, response); err != nil { + t.Errorf("read stream %d: %v", index, err) + return + } + if string(response) != string(request) { + t.Errorf("stream %d response = %q, want %q", index, response, request) + } + }(index, conn) + } + wg.Wait() + if got := carrierDials.Load(); got != 1 { + t.Fatalf("carrier dials = %d, want 1", got) + } +} + +func TestPoolConcurrentStreamStress(t *testing.T) { + var carrierDials atomic.Int32 + pool := NewPool(func(context.Context) (net.Conn, error) { + carrierDials.Add(1) + client, server := net.Pipe() + go serveIndependentEcho(server) + return client, nil + }, Options{ + MaxConcurrency: 4, + MaxConnections: 64, + FirstPayloadTimeout: time.Hour, + IdleTimeout: time.Hour, + }) + t.Cleanup(func() { _ = pool.Close() }) + + const streamCount = 40 + type stream struct { + conn net.Conn + cancel context.CancelFunc + } + streams := make([]stream, 0, streamCount) + for index := 0; index < streamCount; index++ { + ctx, cancel := context.WithCancel(context.Background()) + conn, err := pool.DialContext(ctx, fmt.Sprintf("stress-%d.example", index), uint16(9000+index)) + if err != nil { + t.Fatalf("dial stream %d: %v", index, err) + } + streams = append(streams, stream{conn: conn, cancel: cancel}) + } + + start := make(chan struct{}) + var wg sync.WaitGroup + for index, current := range streams { + wg.Add(1) + go func(index int, current stream) { + defer wg.Done() + <-start + stopDeadlines := make(chan struct{}) + deadlinesDone := make(chan struct{}) + go func() { + defer close(deadlinesDone) + for { + select { + case <-stopDeadlines: + return + default: + _ = current.conn.SetDeadline(time.Now().Add(time.Second)) + _ = current.conn.SetDeadline(time.Time{}) + } + } + }() + + request := []byte(fmt.Sprintf("stress-payload-%d", index)) + if _, err := current.conn.Write(request); err != nil { + t.Errorf("write stream %d: %v", index, err) + } else { + response := make([]byte, len(request)) + if _, err := io.ReadFull(current.conn, response); err != nil { + t.Errorf("read stream %d: %v", index, err) + } else if string(response) != string(request) { + t.Errorf("stream %d response = %q", index, response) + } + } + if index%2 == 0 { + current.cancel() + } else { + _ = current.conn.Close() + } + close(stopDeadlines) + <-deadlinesDone + current.cancel() + _ = current.conn.Close() + }(index, current) + } + close(start) + wg.Wait() + if got := carrierDials.Load(); got < 2 { + t.Fatalf("carrier dials = %d, want several carriers", got) + } + waitFor(t, func() bool { return pool.activeSessions() == 0 }) +} + +func TestPoolInteroperatesWithIndependentXrayCompatiblePacketServer(t *testing.T) { + observed := make(chan independentFrame, 2) + pool := NewPool(func(context.Context) (net.Conn, error) { + client, server := net.Pipe() + go serveIndependentPacketEcho(server, observed) + return client, nil + }, Options{MaxConcurrency: 8, MaxConnections: 128, IdleTimeout: time.Hour}) + t.Cleanup(func() { _ = pool.Close() }) + + packetConn, err := pool.ListenPacketContext(context.Background(), "dns.example", 53, [8]byte{}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packetConn.Close() }) + + requests := []struct { + payload string + addr *net.UDPAddr + want string + }{ + {payload: "one", addr: &net.UDPAddr{IP: net.IPv4(8, 8, 8, 8), Port: 53}, want: "dns.example:53"}, + {payload: "second", addr: &net.UDPAddr{IP: net.IPv4(1, 1, 1, 1), Port: 853}, want: "1.1.1.1:853"}, + } + for _, request := range requests { + if _, err := packetConn.WriteTo([]byte(request.payload), request.addr); err != nil { + t.Fatal(err) + } + buffer := make([]byte, 64) + n, addr, err := packetConn.ReadFrom(buffer) + if err != nil { + t.Fatal(err) + } + if string(buffer[:n]) != request.payload || addr.String() != request.want { + t.Fatalf("packet response = (%q, %v), want (%q, %s)", buffer[:n], addr, request.payload, request.want) + } + } + first := <-observed + second := <-observed + if first.status != 1 || first.host != "dns.example" || second.status != 2 || second.host != "1.1.1.1" { + t.Fatalf("observed frames = first %+v, second %+v", first, second) + } +} + +func TestPoolPreservesXUDPGlobalIDAcrossCarrierRebind(t *testing.T) { + globalID := [8]byte{9, 8, 7, 6, 5, 4, 3, 2} + observed := make(chan independentFrame, 2) + var carrierDials atomic.Int32 + pool := NewPool(func(context.Context) (net.Conn, error) { + carrierDials.Add(1) + client, server := net.Pipe() + go serveIndependentPacketEcho(server, observed) + return client, nil + }, Options{MaxConcurrency: 8, MaxConnections: 1, IdleTimeout: time.Hour}) + t.Cleanup(func() { _ = pool.Close() }) + + for index := 0; index < 2; index++ { + packetConn, err := pool.ListenPacketContext(context.Background(), "rebind.example", 443, globalID) + if err != nil { + t.Fatal(err) + } + payload := []byte{byte('a' + index)} + if _, err := packetConn.WriteTo(payload, &net.UDPAddr{IP: net.IPv4(203, 0, 113, 1), Port: 443}); err != nil { + t.Fatal(err) + } + buffer := make([]byte, 1) + if _, _, err := packetConn.ReadFrom(buffer); err != nil { + t.Fatal(err) + } + if err := packetConn.Close(); err != nil { + t.Fatal(err) + } + waitFor(t, func() bool { return pool.activeSessions() == 0 }) + waitFor(t, func() bool { return pool.workerCount() == 0 }) + } + + first := <-observed + second := <-observed + if first.globalID != globalID || second.globalID != globalID { + t.Fatalf("rebind GlobalIDs = %v, %v", first.globalID, second.globalID) + } + if got := carrierDials.Load(); got != 2 { + t.Fatalf("carrier dials = %d, want 2", got) + } +} diff --git a/transport/muxcool/packet_session.go b/transport/muxcool/packet_session.go new file mode 100644 index 0000000000..4bf36c7d48 --- /dev/null +++ b/transport/muxcool/packet_session.go @@ -0,0 +1,355 @@ +package muxcool + +import ( + "errors" + "fmt" + "net" + "net/netip" + "os" + "strconv" + "sync" + "sync/atomic" + "time" + + "github.com/metacubex/mihomo/common/net/deadline" + "github.com/metacubex/mihomo/common/pool" +) + +var ErrPacketTooLarge = fmt.Errorf("mux.cool packet exceeds %d bytes", MaxPayloadSize) + +type packetMessage struct { + payload []byte + addr net.Addr +} + +func consumePacketMessage(p []byte, message packetMessage) (int, net.Addr, error) { + n := copy(p, message.payload) + _ = pool.Put(message.payload) + return n, message.addr, nil +} + +type packetSession struct { + owner sessionOwner + id uint16 + destination string + port uint16 + globalID [8]byte + + writeMu sync.Mutex + sentNew bool + + input chan packetMessage + done chan struct{} + closeOnce sync.Once + closedFast atomic.Bool + causeMu sync.Mutex + cause error + readDeadline deadline.PipeDeadline + writeDeadline deadline.PipeDeadline + writeDeadlineSet atomic.Bool +} + +func newPacketSession( + owner sessionOwner, + id uint16, + destination string, + port uint16, + globalID [8]byte, +) (net.PacketConn, *packetSession) { + s := makePacketSession(owner, id, destination, port, globalID) + return s, s +} + +func makePacketSession( + owner sessionOwner, + id uint16, + destination string, + port uint16, + globalID [8]byte, +) *packetSession { + s := &packetSession{ + owner: owner, + id: id, + destination: destination, + port: port, + globalID: globalID, + input: make(chan packetMessage, 16), + done: make(chan struct{}), + readDeadline: deadline.MakePipeDeadline(), + writeDeadline: deadline.MakePipeDeadline(), + } + return s +} + +func (s *packetSession) ReadFrom(p []byte) (int, net.Addr, error) { + message, err := s.readPacketMessage() + if err != nil { + return 0, nil, err + } + return consumePacketMessage(p, message) +} + +func (s *packetSession) WaitReadFrom() ([]byte, func(), net.Addr, error) { + message, err := s.readPacketMessage() + if err != nil { + return nil, nil, nil, err + } + payload := message.payload + put := func() { + _ = pool.Put(payload) + } + return payload, put, message.addr, nil +} + +func (s *packetSession) readPacketMessage() (packetMessage, error) { + if message, ok := s.nextQueued(); ok { + return message, nil + } + select { + case message := <-s.input: + return message, nil + case <-s.done: + if message, ok := s.nextQueued(); ok { + return message, nil + } + return packetMessage{}, s.terminalCause() + case <-s.readDeadline.Wait(): + return packetMessage{}, os.ErrDeadlineExceeded + } +} + +func (s *packetSession) nextQueued() (packetMessage, bool) { + select { + case message := <-s.input: + return message, true + default: + return packetMessage{}, false + } +} + +func (s *packetSession) WriteTo(payload []byte, addr net.Addr) (int, error) { + if len(payload) == 0 { + return 0, nil + } + if len(payload) > MaxPayloadSize { + return 0, ErrPacketTooLarge + } + target, err := splitPacketAddr(addr) + if err != nil { + return 0, err + } + + s.writeMu.Lock() + if s.closedFast.Load() { + s.writeMu.Unlock() + return 0, s.terminalCause() + } + if s.writeDeadlineSet.Load() { + select { + case <-s.writeDeadline.Wait(): + s.writeMu.Unlock() + return 0, os.ErrDeadlineExceeded + default: + } + } + + frame := Frame{ + SessionID: s.id, + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + Destination: target.host, + DestinationIP: target.ip, + Port: target.port, + Payload: payload, + } + if !s.sentNew { + frame.Status = StatusNew + frame.Destination = s.destination + frame.DestinationIP = netip.Addr{} + frame.Port = s.port + frame.GlobalID = s.globalID + } + err = s.owner.writeFrame(frame) + if err == nil { + s.sentNew = true + } + s.writeMu.Unlock() + if err != nil { + s.finish(err, false) + return 0, err + } + return len(payload), nil +} + +func (s *packetSession) Close() error { + s.finish(nil, true) + return nil +} + +func (s *packetSession) LocalAddr() net.Addr { + return muxAddr("mux.cool-udp") +} + +func (s *packetSession) SetDeadline(t time.Time) error { + s.readDeadline.Set(t) + return s.SetWriteDeadline(t) +} + +func (s *packetSession) SetReadDeadline(t time.Time) error { + s.readDeadline.Set(t) + return nil +} + +func (s *packetSession) SetWriteDeadline(t time.Time) error { + if !t.IsZero() { + s.writeDeadlineSet.Store(true) + } + s.writeDeadline.Set(t) + if t.IsZero() { + s.writeDeadlineSet.Store(false) + } + return nil +} + +func (s *packetSession) deliverFrame(frame Frame) error { + decoded := decodedFrame{Frame: frame} + if len(frame.Payload) > 0 { + decoded.Payload = pool.Get(len(frame.Payload))[:len(frame.Payload)] + copy(decoded.Payload, frame.Payload) + decoded.payloadPooled = true + } + return s.deliverDecodedFrame(decoded) +} + +func (s *packetSession) deliverDecodedFrame(decoded decodedFrame) error { + frame := decoded.Frame + if frame.Option&OptionData != 0 { + if frame.Network != NetworkUDP || frame.Destination == "" && !frame.DestinationIP.IsValid() { + decoded.releasePayload() + return protocolError("deliver UDP", errors.New("response frame is missing a UDP target")) + } + addr, err := makeDecodedPacketAddr(frame) + if err != nil { + decoded.releasePayload() + return protocolError("deliver UDP", err) + } + select { + case s.input <- packetMessage{payload: frame.Payload, addr: addr}: + case <-s.done: + decoded.releasePayload() + return net.ErrClosed + } + } else { + decoded.releasePayload() + } + if frame.Status == StatusEnd || frame.Option&OptionError != 0 { + var cause error + if frame.Option&OptionError != 0 { + cause = protocolError("remote session", errors.New("remote reported an error")) + } + s.finish(cause, false) + } + return nil +} + +func makeDecodedPacketAddr(frame Frame) (net.Addr, error) { + if frame.DestinationIP.IsValid() { + if frame.DestinationIP.Zone() != "" { + return nil, errors.New("scoped IP addresses are not supported") + } + return net.UDPAddrFromAddrPort(netip.AddrPortFrom(frame.DestinationIP.Unmap(), frame.Port)), nil + } + return makePacketAddr(frame.Destination, frame.Port) +} + +func (s *packetSession) closeCarrier(cause error) { + s.finish(cause, false) +} + +func (s *packetSession) finish(cause error, sendEnd bool) { + s.closeOnce.Do(func() { + s.writeMu.Lock() + if sendEnd { + _ = s.owner.writeFrame(Frame{SessionID: s.id, Status: StatusEnd}) + } + s.causeMu.Lock() + s.cause = cause + s.causeMu.Unlock() + s.closedFast.Store(true) + close(s.done) + s.writeMu.Unlock() + s.owner.removeSession(s.id) + }) +} + +func (s *packetSession) terminalCause() error { + s.causeMu.Lock() + defer s.causeMu.Unlock() + if s.cause != nil { + return s.cause + } + return net.ErrClosed +} + +type domainPacketAddr struct { + host string + port uint16 +} + +func (a domainPacketAddr) Network() string { return "udp" } +func (a domainPacketAddr) String() string { + return net.JoinHostPort(a.host, strconv.Itoa(int(a.port))) +} + +func makePacketAddr(host string, port uint16) (net.Addr, error) { + if ip, err := netip.ParseAddr(host); err == nil { + if ip.Zone() != "" { + return nil, errors.New("scoped IPv6 addresses are not supported") + } + return net.UDPAddrFromAddrPort(netip.AddrPortFrom(ip.Unmap(), port)), nil + } + if host == "" { + return nil, errors.New("empty packet host") + } + return domainPacketAddr{host: host, port: port}, nil +} + +type packetTarget struct { + host string + ip netip.Addr + port uint16 +} + +func splitPacketAddr(addr net.Addr) (packetTarget, error) { + if addr == nil { + return packetTarget{}, errors.New("nil packet address") + } + if udpAddr, ok := addr.(*net.UDPAddr); ok { + if udpAddr.IP == nil || udpAddr.Port < 0 || udpAddr.Port > int(^uint16(0)) || udpAddr.Zone != "" { + return packetTarget{}, fmt.Errorf("invalid UDP address %v", addr) + } + ip, valid := netip.AddrFromSlice(udpAddr.IP) + if !valid { + return packetTarget{}, fmt.Errorf("invalid UDP address %v", addr) + } + return packetTarget{ip: ip.Unmap(), port: uint16(udpAddr.Port)}, nil + } + host, rawPort, err := net.SplitHostPort(addr.String()) + if err != nil { + return packetTarget{}, fmt.Errorf("parse packet address %q: %w", addr.String(), err) + } + port, err := strconv.ParseUint(rawPort, 10, 16) + if err != nil || host == "" { + return packetTarget{}, fmt.Errorf("invalid packet address %q", addr.String()) + } + if ip, parseErr := netip.ParseAddr(host); parseErr == nil { + if ip.Zone() != "" { + return packetTarget{}, errors.New("scoped IPv6 addresses are not supported") + } + return packetTarget{ip: ip.Unmap(), port: uint16(port)}, nil + } + return packetTarget{host: host, port: uint16(port)}, nil +} + +var _ net.PacketConn = (*packetSession)(nil) diff --git a/transport/muxcool/packet_session_test.go b/transport/muxcool/packet_session_test.go new file mode 100644 index 0000000000..ba58d73606 --- /dev/null +++ b/transport/muxcool/packet_session_test.go @@ -0,0 +1,131 @@ +package muxcool + +import ( + "net" + "net/netip" + "testing" + "time" +) + +func TestPacketSessionWritesXrayUDPFrames(t *testing.T) { + owner := newFakeSessionOwner() + globalID := [8]byte{1, 2, 3, 4, 5, 6, 7, 8} + packetConn, _ := newPacketSession(owner, 31, "initial.example", 53, globalID) + t.Cleanup(func() { _ = packetConn.Close() }) + + firstAddr := net.UDPAddrFromAddrPort(netip.MustParseAddrPort("8.8.8.8:53")) + if n, err := packetConn.WriteTo([]byte("first"), firstAddr); err != nil || n != 5 { + t.Fatalf("first WriteTo = (%d, %v)", n, err) + } + first := receiveFrame(t, owner.frames) + if first.Status != StatusNew || first.Network != NetworkUDP || first.Destination != "initial.example" || first.Port != 53 { + t.Fatalf("first frame target = %+v", first) + } + if first.GlobalID != globalID || string(first.Payload) != "first" { + t.Fatalf("first frame = %+v", first) + } + + secondAddr := net.UDPAddrFromAddrPort(netip.MustParseAddrPort("[2001:db8::1]:5353")) + if n, err := packetConn.WriteTo([]byte("second"), secondAddr); err != nil || n != 6 { + t.Fatalf("second WriteTo = (%d, %v)", n, err) + } + second := receiveFrame(t, owner.frames) + if second.Status != StatusKeep || second.Network != NetworkUDP || second.DestinationIP.String() != "2001:db8::1" || second.Port != 5353 { + t.Fatalf("second frame target = %+v", second) + } + if second.GlobalID != [8]byte{} || string(second.Payload) != "second" { + t.Fatalf("second frame = %+v", second) + } +} + +func TestPacketSessionPreservesDatagramBoundariesAndAddresses(t *testing.T) { + owner := newFakeSessionOwner() + packetConn, session := newPacketSession(owner, 32, "initial.example", 53, [8]byte{}) + t.Cleanup(func() { _ = packetConn.Close() }) + + if err := session.deliverFrame(Frame{ + SessionID: 32, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, + Destination: "1.1.1.1", Port: 853, Payload: []byte("first-packet"), + }); err != nil { + t.Fatal(err) + } + if err := session.deliverFrame(Frame{ + SessionID: 32, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, + Destination: "reply.example", Port: 5353, Payload: []byte("two"), + }); err != nil { + t.Fatal(err) + } + + buffer := make([]byte, 5) + n, addr, err := packetConn.ReadFrom(buffer) + if err != nil { + t.Fatal(err) + } + if n != len(buffer) || string(buffer) != "first" || addr.String() != "1.1.1.1:853" { + t.Fatalf("first ReadFrom = (%d, %q, %v)", n, buffer, addr) + } + + buffer = make([]byte, 16) + n, addr, err = packetConn.ReadFrom(buffer) + if err != nil { + t.Fatal(err) + } + if n != 3 || string(buffer[:n]) != "two" || addr.String() != "reply.example:5353" { + t.Fatalf("second ReadFrom = (%d, %q, %v)", n, buffer[:n], addr) + } +} + +func TestPacketSessionDeadlinesAndExplicitClose(t *testing.T) { + owner := newFakeSessionOwner() + packetConn, _ := newPacketSession(owner, 33, "cancel.example", 53, [8]byte{}) + + if err := packetConn.SetReadDeadline(time.Now().Add(-time.Second)); err != nil { + t.Fatal(err) + } + if _, _, err := packetConn.ReadFrom(make([]byte, 1)); !isTimeout(err) { + t.Fatalf("ReadFrom error = %v, want timeout", err) + } + if err := packetConn.SetReadDeadline(time.Time{}); err != nil { + t.Fatal(err) + } + if err := packetConn.SetWriteDeadline(time.Now().Add(-time.Second)); err != nil { + t.Fatal(err) + } + if _, err := packetConn.WriteTo([]byte("x"), net.UDPAddrFromAddrPort(netip.MustParseAddrPort("1.1.1.1:53"))); !isTimeout(err) { + t.Fatalf("WriteTo error = %v, want timeout", err) + } + if err := packetConn.SetDeadline(time.Time{}); err != nil { + t.Fatal(err) + } + + if n, err := packetConn.WriteTo([]byte("alive"), net.UDPAddrFromAddrPort(netip.MustParseAddrPort("1.1.1.1:53"))); err != nil || n != len("alive") { + t.Fatalf("WriteTo after clearing deadline = (%d, %v)", n, err) + } + data := receiveFrame(t, owner.frames) + if data.Status != StatusNew || string(data.Payload) != "alive" { + t.Fatalf("data after clearing deadline = %+v", data) + } + if err := packetConn.Close(); err != nil { + t.Fatal(err) + } + if err := packetConn.Close(); err != nil { + t.Fatal(err) + } + end := receiveFrame(t, owner.frames) + if end.Status != StatusEnd || end.SessionID != 33 { + t.Fatalf("End = %+v", end) + } + select { + case id := <-owner.removed: + if id != 33 { + t.Fatalf("removed ID = %d", id) + } + case <-time.After(time.Second): + t.Fatal("packet session was not removed on Close") + } + select { + case extra := <-owner.frames: + t.Fatalf("unexpected extra frame: %+v", extra) + case <-time.After(20 * time.Millisecond): + } +} diff --git a/transport/muxcool/performance_test.go b/transport/muxcool/performance_test.go new file mode 100644 index 0000000000..4c8fe4d598 --- /dev/null +++ b/transport/muxcool/performance_test.go @@ -0,0 +1,608 @@ +package muxcool + +import ( + "bytes" + "context" + "io" + "net" + "net/netip" + "testing" + "time" + + N "github.com/metacubex/mihomo/common/net" + "github.com/metacubex/mihomo/common/net/deadline" + "github.com/metacubex/mihomo/common/pool" + "github.com/metacubex/sing/common/buf" + M "github.com/metacubex/sing/common/metadata" +) + +var ( + benchmarkBytes []byte + benchmarkFrame Frame + benchmarkInt int +) + +func BenchmarkEncodeFrame(b *testing.B) { + benchmarks := []struct { + name string + frame Frame + }{ + { + name: "new-domain-8k", + frame: Frame{ + SessionID: 1, + Status: StatusNew, + Option: OptionData, + Network: NetworkTCP, + Destination: "example.com", + Port: 443, + Payload: make([]byte, MaxPayloadSize), + }, + }, + { + name: "keep-8k", + frame: Frame{ + SessionID: 1, + Status: StatusKeep, + Option: OptionData, + Payload: make([]byte, MaxPayloadSize), + }, + }, + { + name: "udp-new-domain-1k", + frame: Frame{ + SessionID: 2, + Status: StatusNew, + Option: OptionData, + Network: NetworkUDP, + Destination: "dns.example", + Port: 53, + GlobalID: [8]byte{1, 2, 3, 4, 5, 6, 7, 8}, + Payload: make([]byte, 1024), + }, + }, + { + name: "udp-keep-ipv4-1k", + frame: Frame{ + SessionID: 2, + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + Destination: "1.1.1.1", + Port: 53, + Payload: make([]byte, 1024), + }, + }, + { + name: "udp-keep-netip-1k", + frame: Frame{ + SessionID: 2, + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + DestinationIP: netip.MustParseAddr("1.1.1.1"), + Port: 53, + Payload: make([]byte, 1024), + }, + }, + } + + for _, benchmark := range benchmarks { + b.Run(benchmark.name, func(b *testing.B) { + b.ReportAllocs() + b.SetBytes(int64(len(benchmark.frame.Payload))) + for i := 0; i < b.N; i++ { + encoded, err := EncodeFrame(benchmark.frame) + if err != nil { + b.Fatal(err) + } + benchmarkBytes = encoded + } + }) + } +} + +func BenchmarkDecodeFrame(b *testing.B) { + raw, err := EncodeFrame(Frame{ + SessionID: 1, + Status: StatusKeep, + Option: OptionData, + Payload: make([]byte, MaxPayloadSize), + }) + if err != nil { + b.Fatal(err) + } + + var reader bytes.Reader + metadataBuffer := make([]byte, MaxMetadataSize) + b.ReportAllocs() + b.SetBytes(MaxPayloadSize) + for i := 0; i < b.N; i++ { + reader.Reset(raw) + frame, err := decodeFrame(&reader, metadataBuffer) + if err != nil { + b.Fatal(err) + } + benchmarkFrame = frame + } +} + +func BenchmarkDecodeFramePooled(b *testing.B) { + raw, err := EncodeFrame(Frame{ + SessionID: 1, + Status: StatusKeep, + Option: OptionData, + Payload: make([]byte, MaxPayloadSize), + }) + if err != nil { + b.Fatal(err) + } + + var reader bytes.Reader + metadataBuffer := make([]byte, MaxMetadataSize) + b.ReportAllocs() + b.SetBytes(MaxPayloadSize) + for i := 0; i < b.N; i++ { + reader.Reset(raw) + frame, err := decodeFramePooled(&reader, metadataBuffer) + if err != nil { + b.Fatal(err) + } + benchmarkInt = len(frame.Payload) + frame.releasePayload() + } +} + +func BenchmarkDecodeUDPFrame(b *testing.B) { + raw, err := EncodeFrame(Frame{ + SessionID: 2, + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + Destination: "1.1.1.1", + Port: 53, + Payload: make([]byte, 1024), + }) + if err != nil { + b.Fatal(err) + } + + var reader bytes.Reader + metadataBuffer := make([]byte, MaxMetadataSize) + b.ReportAllocs() + b.SetBytes(1024) + for i := 0; i < b.N; i++ { + reader.Reset(raw) + frame, err := decodeFrame(&reader, metadataBuffer) + if err != nil { + b.Fatal(err) + } + benchmarkFrame = frame + } +} + +func BenchmarkSessionDeliver(b *testing.B) { + s := &session{ + downlink: make(chan downlinkMessage, 1), + done: make(chan struct{}), + } + payload := make([]byte, MaxPayloadSize) + + b.ReportAllocs() + b.SetBytes(MaxPayloadSize) + for i := 0; i < b.N; i++ { + if err := s.deliver(payload); err != nil { + b.Fatal(err) + } + message := <-s.downlink + benchmarkInt = len(message.payload) + } +} + +func BenchmarkDownlinkMessageRelease(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + message := downlinkMessage{payload: pool.Get(1024)[:1024], payloadPooled: true} + message.releasePayload() + } +} + +func BenchmarkWriteStreamData(b *testing.B) { + payload := make([]byte, 64*1024) + b.ReportAllocs() + b.SetBytes(int64(len(payload))) + for i := 0; i < b.N; i++ { + if err := writeStreamData(io.Discard, 1, "example.com", 443, true, payload); err != nil { + b.Fatal(err) + } + } +} + +type benchmarkSessionOwner struct { + written chan struct{} + buffer []byte +} + +func (o *benchmarkSessionOwner) writeFrame(frame Frame) error { + encoded, err := encodeFrame(o.buffer, frame) + if err != nil { + return err + } + o.buffer = encoded[:0] + benchmarkBytes = encoded + if frame.Option&OptionData != 0 { + o.written <- struct{}{} + } + return nil +} + +func (*benchmarkSessionOwner) removeSession(uint16) {} + +type benchmarkNoopOwner struct{} + +func (benchmarkNoopOwner) writeFrame(Frame) error { return nil } +func (benchmarkNoopOwner) removeSession(uint16) {} + +func BenchmarkPacketSessionLifecycle(b *testing.B) { + owner := benchmarkNoopOwner{} + b.ReportAllocs() + for i := 0; i < b.N; i++ { + session := makePacketSession(owner, 2, "dns.example", 53, [8]byte{}) + _ = session.Close() + } +} + +type benchmarkConn struct{} + +func (benchmarkConn) Read([]byte) (int, error) { return 0, io.EOF } +func (benchmarkConn) Write(p []byte) (int, error) { return len(p), nil } +func (benchmarkConn) Close() error { return nil } +func (benchmarkConn) LocalAddr() net.Addr { return muxAddr("local") } +func (benchmarkConn) RemoteAddr() net.Addr { return muxAddr("remote") } +func (benchmarkConn) SetDeadline(time.Time) error { return nil } +func (benchmarkConn) SetReadDeadline(time.Time) error { return nil } +func (benchmarkConn) SetWriteDeadline(time.Time) error { return nil } + +func BenchmarkCarrierWorkerWriteFrame(b *testing.B) { + worker := &carrierWorker{conn: benchmarkConn{}} + frame := Frame{SessionID: 1, Status: StatusKeep, Option: OptionData, Payload: make([]byte, MaxPayloadSize)} + if err := worker.writeFrame(frame); err != nil { + b.Fatal(err) + } + b.ReportAllocs() + b.SetBytes(MaxPayloadSize) + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := worker.writeFrame(frame); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkServerCarrierWriteFrame(b *testing.B) { + carrier := &serverCarrier{conn: benchmarkConn{}} + frame := Frame{SessionID: 1, Status: StatusKeep, Option: OptionData, Payload: make([]byte, MaxPayloadSize)} + if err := carrier.writeFrame(frame); err != nil { + b.Fatal(err) + } + b.ReportAllocs() + b.SetBytes(MaxPayloadSize) + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := carrier.writeFrame(frame); err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCarrierWorkerClose(b *testing.B) { + b.ReportAllocs() + for i := 0; i < b.N; i++ { + worker := &carrierWorker{ + conn: benchmarkConn{}, + sessions: make(map[uint16]workerSession), + } + worker.close(nil) + } +} + +type benchmarkTimer struct{} + +func (benchmarkTimer) Stop() bool { return true } + +func BenchmarkPoolPacketSessionChurn(b *testing.B) { + servers := make([]net.Conn, 0, 1) + pool := NewPool(func(context.Context) (net.Conn, error) { + client, server := net.Pipe() + servers = append(servers, server) + go func() { _, _ = io.Copy(io.Discard, server) }() + return client, nil + }, Options{ + MaxConcurrency: 1, + MaxConnections: int(^uint(0) >> 1), + AfterFunc: func(time.Duration, func()) Timer { + return benchmarkTimer{} + }, + }) + b.Cleanup(func() { + _ = pool.Close() + for _, server := range servers { + _ = server.Close() + } + }) + + warm, err := pool.ListenPacketContext(context.Background(), "dns.example", 53, [8]byte{}) + if err != nil { + b.Fatal(err) + } + _ = warm.Close() + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + packetConn, err := pool.ListenPacketContext(context.Background(), "dns.example", 53, [8]byte{}) + if err != nil { + b.Fatal(err) + } + _ = packetConn.Close() + } +} + +func BenchmarkPoolPacketSessionChurnParallel(b *testing.B) { + var server net.Conn + pool := NewPool(func(context.Context) (net.Conn, error) { + client, peer := net.Pipe() + server = peer + go func() { _, _ = io.Copy(io.Discard, peer) }() + return client, nil + }, Options{ + MaxConcurrency: 1 << 20, + MaxConnections: int(^uint(0) >> 1), + AfterFunc: func(time.Duration, func()) Timer { + return benchmarkTimer{} + }, + }) + b.Cleanup(func() { + _ = pool.Close() + if server != nil { + _ = server.Close() + } + }) + + warm, err := pool.ListenPacketContext(context.Background(), "dns.example", 53, [8]byte{}) + if err != nil { + b.Fatal(err) + } + _ = warm.Close() + b.ReportAllocs() + b.RunParallel(func(pb *testing.PB) { + for pb.Next() { + packetConn, err := pool.ListenPacketContext(context.Background(), "dns.example", 53, [8]byte{}) + if err != nil { + b.Error(err) + return + } + _ = packetConn.Close() + } + }) +} + +func BenchmarkStreamSessionLifecycle(b *testing.B) { + owner := benchmarkNoopOwner{} + payload := []byte{1} + b.ReportAllocs() + for i := 0; i < b.N; i++ { + conn, session := makeSession(owner, 1, "example.com", 443) + session.start(context.Background(), 0) + if _, err := conn.Write(payload); err != nil { + b.Fatal(err) + } + _ = conn.Close() + <-session.done + } +} + +func BenchmarkSessionUplinkFrame(b *testing.B) { + owner := &benchmarkSessionOwner{written: make(chan struct{})} + conn, session := makeSession(owner, 1, "example.com", 443) + session.start(context.Background(), 0) + payload := make([]byte, MaxPayloadSize) + b.Cleanup(func() { _ = conn.Close() }) + + b.ReportAllocs() + b.SetBytes(MaxPayloadSize) + b.ResetTimer() + for i := 0; i < b.N; i++ { + if _, err := conn.Write(payload); err != nil { + b.Fatal(err) + } + <-owner.written + } +} + +func BenchmarkPacketSessionWriteTo(b *testing.B) { + owner := &benchmarkSessionOwner{written: make(chan struct{}, 1)} + session := makePacketSession(owner, 2, "dns.example", 53, [8]byte{1, 2, 3, 4, 5, 6, 7, 8}) + addr := &net.UDPAddr{IP: net.IPv4(1, 1, 1, 1), Port: 53} + payload := make([]byte, 1024) + if _, err := session.WriteTo(payload, addr); err != nil { + b.Fatal(err) + } + <-owner.written + + b.ReportAllocs() + b.SetBytes(int64(len(payload))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + if _, err := session.WriteTo(payload, addr); err != nil { + b.Fatal(err) + } + <-owner.written + } +} + +func BenchmarkPacketSessionDeliverReadFrom(b *testing.B) { + session := makePacketSession(&benchmarkSessionOwner{}, 2, "dns.example", 53, [8]byte{}) + frame := Frame{ + SessionID: 2, + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + Destination: "1.1.1.1", + Port: 53, + Payload: make([]byte, 1024), + } + buffer := make([]byte, len(frame.Payload)) + + b.ReportAllocs() + b.SetBytes(int64(len(frame.Payload))) + b.ResetTimer() + for i := 0; i < b.N; i++ { + if err := session.deliverFrame(frame); err != nil { + b.Fatal(err) + } + n, _, err := session.ReadFrom(buffer) + if err != nil { + b.Fatal(err) + } + benchmarkInt = n + } +} + +func BenchmarkPacketSessionDecodeDeliverWaitReadFrom(b *testing.B) { + raw, err := EncodeFrame(Frame{ + SessionID: 2, + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + Destination: "1.1.1.1", + Port: 53, + Payload: make([]byte, 1024), + }) + if err != nil { + b.Fatal(err) + } + session := makePacketSession(&benchmarkSessionOwner{}, 2, "dns.example", 53, [8]byte{}) + var reader bytes.Reader + metadataBuffer := make([]byte, MaxMetadataSize) + + b.ReportAllocs() + b.SetBytes(1024) + b.ResetTimer() + for i := 0; i < b.N; i++ { + reader.Reset(raw) + frame, err := decodeFramePooled(&reader, metadataBuffer) + if err != nil { + b.Fatal(err) + } + if err := session.deliverDecodedFrame(frame); err != nil { + b.Fatal(err) + } + data, put, _, err := session.WaitReadFrom() + if err != nil { + b.Fatal(err) + } + benchmarkInt = len(data) + if put != nil { + put() + } + } +} + +type benchmarkPacketConnOnly struct { + net.PacketConn +} + +func BenchmarkPacketSessionDecodeDeliverCopiedWaitReadFrom(b *testing.B) { + raw, err := EncodeFrame(Frame{ + SessionID: 2, + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + Destination: "1.1.1.1", + Port: 53, + Payload: make([]byte, 1024), + }) + if err != nil { + b.Fatal(err) + } + session := makePacketSession(&benchmarkSessionOwner{}, 2, "dns.example", 53, [8]byte{}) + packetConn := N.NewEnhancePacketConn(benchmarkPacketConnOnly{PacketConn: session}) + var reader bytes.Reader + metadataBuffer := make([]byte, MaxMetadataSize) + + b.ReportAllocs() + b.SetBytes(1024) + b.ResetTimer() + for i := 0; i < b.N; i++ { + reader.Reset(raw) + frame, err := decodeFramePooled(&reader, metadataBuffer) + if err != nil { + b.Fatal(err) + } + if err := session.deliverDecodedFrame(frame); err != nil { + b.Fatal(err) + } + data, put, _, err := packetConn.WaitReadFrom() + if err != nil { + b.Fatal(err) + } + benchmarkInt = len(data) + if put != nil { + put() + } + } +} + +func BenchmarkServerPacketFlowEnqueue(b *testing.B) { + flow := &serverPacketFlow{ + input: make(chan serverPacketMessage, 1), + done: make(chan struct{}), + } + attachment := &serverPacketAttachment{flow: flow, generation: 1} + flow.current = attachment + flow.currentFast.Store(attachment) + flow.generation = 1 + frame := decodedFrame{Frame: Frame{ + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + DestinationIP: netip.MustParseAddr("1.1.1.1"), + Port: 53, + Payload: make([]byte, 1024), + }} + + b.ReportAllocs() + b.SetBytes(1024) + for i := 0; i < b.N; i++ { + if err := flow.enqueue(attachment, frame); err != nil { + b.Fatal(err) + } + message := <-flow.input + benchmarkInt = len(message.payload) + } +} + +func BenchmarkServerPacketFlowWritePacket(b *testing.B) { + carrier := &serverCarrier{conn: benchmarkConn{}} + flow := &serverPacketFlow{ + done: make(chan struct{}), + writeDeadline: deadline.MakePipeDeadline(), + } + attachment := &serverPacketAttachment{flow: flow, carrier: carrier, id: 1, generation: 1} + flow.current = attachment + flow.currentFast.Store(attachment) + packet := buf.NewSize(1024) + if _, err := packet.Write(make([]byte, 1024)); err != nil { + b.Fatal(err) + } + b.Cleanup(packet.Release) + destination := M.Socksaddr{Addr: netip.MustParseAddr("1.1.1.1"), Port: 53} + + b.ReportAllocs() + b.SetBytes(1024) + for i := 0; i < b.N; i++ { + if err := flow.WritePacket(packet, destination); err != nil { + b.Fatal(err) + } + } +} diff --git a/transport/muxcool/pool.go b/transport/muxcool/pool.go new file mode 100644 index 0000000000..7e17db2389 --- /dev/null +++ b/transport/muxcool/pool.go @@ -0,0 +1,332 @@ +package muxcool + +import ( + "context" + "errors" + "net" + "sync" + "time" +) + +const ( + DefaultMaxConcurrency = 8 + DefaultMaxConnections = 128 + DefaultFirstPayloadTimeout = 100 * time.Millisecond + DefaultIdleTimeout = 16 * time.Second +) + +var ErrPoolClosed = errors.New("mux.cool pool is closed") + +type Timer interface { + Stop() bool +} + +type CarrierDialer func(context.Context) (net.Conn, error) + +type Options struct { + MaxConcurrency int + MaxConnections int + CarrierLimiter *CarrierLimiter + FirstPayloadTimeout time.Duration + IdleTimeout time.Duration + AfterFunc func(time.Duration, func()) Timer +} + +type Pool struct { + dial CarrierDialer + options Options + + mu sync.Mutex + workers []*carrierWorker + idle map[*carrierWorker]Timer + dialing chan struct{} + dialCancel context.CancelFunc + changed chan struct{} + done chan struct{} + closed bool + closeOnce sync.Once +} + +func NewPool(dial CarrierDialer, options Options) *Pool { + if options.MaxConcurrency == 0 { + options.MaxConcurrency = DefaultMaxConcurrency + } + if options.MaxConnections == 0 { + options.MaxConnections = DefaultMaxConnections + } + if options.FirstPayloadTimeout == 0 { + options.FirstPayloadTimeout = DefaultFirstPayloadTimeout + } + if options.IdleTimeout == 0 { + options.IdleTimeout = DefaultIdleTimeout + } + if options.AfterFunc == nil { + options.AfterFunc = func(duration time.Duration, fn func()) Timer { + return time.AfterFunc(duration, fn) + } + } + return &Pool{ + dial: dial, + options: options, + idle: make(map[*carrierWorker]Timer), + changed: make(chan struct{}), + done: make(chan struct{}), + } +} + +func (p *Pool) DialContext(ctx context.Context, destination string, port uint16) (net.Conn, error) { + opened, err := p.openContext(ctx, destination, port, [8]byte{}, false) + if err != nil { + return nil, err + } + return opened.stream, nil +} + +func (p *Pool) ListenPacketContext(ctx context.Context, destination string, port uint16, globalID [8]byte) (net.PacketConn, error) { + opened, err := p.openContext(ctx, destination, port, globalID, true) + if err != nil { + return nil, err + } + return opened.packet, nil +} + +type openedSession struct { + stream net.Conn + packet net.PacketConn +} + +func (p *Pool) openContext( + ctx context.Context, + destination string, + port uint16, + globalID [8]byte, + packet bool, +) (openedSession, error) { + for { + if err := context.Cause(ctx); err != nil { + return openedSession{}, err + } + p.mu.Lock() + if p.closed { + p.mu.Unlock() + return openedSession{}, ErrPoolClosed + } + + for _, worker := range p.workers { + conn, err := openWorkerSession(worker, ctx, destination, port, globalID, packet, p.options.FirstPayloadTimeout) + if err != nil { + continue + } + p.stopIdleLocked(worker) + p.mu.Unlock() + return conn, nil + } + + if dialing := p.dialing; dialing != nil { + done := p.done + p.mu.Unlock() + select { + case <-dialing: + continue + case <-done: + return openedSession{}, ErrPoolClosed + case <-ctx.Done(): + return openedSession{}, context.Cause(ctx) + } + } + + var lease *carrierLease + if p.options.CarrierLimiter.limited() { + var capacityWait <-chan struct{} + lease, capacityWait = p.options.CarrierLimiter.tryAcquire() + if lease == nil { + changed := p.changed + done := p.done + p.mu.Unlock() + select { + case <-changed: + continue + case <-capacityWait: + continue + case <-done: + return openedSession{}, ErrPoolClosed + case <-ctx.Done(): + return openedSession{}, context.Cause(ctx) + } + } + } + + dialing := make(chan struct{}) + dialCtx, cancel := context.WithCancel(ctx) + p.dialing = dialing + p.dialCancel = cancel + p.mu.Unlock() + + carrier, err := p.dial(dialCtx) + cancel() + if err != nil { + if carrier != nil { + _ = carrier.Close() + } + lease.release() + } else if lease != nil { + carrier = &limitedCarrier{Conn: carrier, lease: lease} + } + + p.mu.Lock() + if p.dialing == dialing { + p.dialing = nil + p.dialCancel = nil + close(dialing) + } + if p.closed { + p.mu.Unlock() + if carrier != nil { + _ = carrier.Close() + } + return openedSession{}, ErrPoolClosed + } + if err != nil { + p.mu.Unlock() + return openedSession{}, err + } + + worker := newCarrierWorker(carrier, p.options.MaxConcurrency, p.options.MaxConnections, p.removeWorker) + worker.onIdle = p.scheduleIdle + if p.options.CarrierLimiter.limited() { + worker.onAvailable = p.signalAvailable + } + p.workers = append(p.workers, worker) + conn, err := openWorkerSession(worker, ctx, destination, port, globalID, packet, p.options.FirstPayloadTimeout) + if err == nil { + p.mu.Unlock() + return conn, nil + } + p.removeWorkerLocked(worker) + p.mu.Unlock() + worker.close(err) + return openedSession{}, err + } +} + +func (p *Pool) signalAvailable(*carrierWorker) { + p.mu.Lock() + p.signalChangedLocked() + p.mu.Unlock() +} + +func (p *Pool) signalChangedLocked() { + changed := p.changed + p.changed = make(chan struct{}) + close(changed) +} + +func openWorkerSession( + worker *carrierWorker, + ctx context.Context, + destination string, + port uint16, + globalID [8]byte, + packet bool, + firstPayloadTimeout time.Duration, +) (openedSession, error) { + if packet { + connection, err := worker.openPacketSession(destination, port, globalID) + return openedSession{packet: connection}, err + } + connection, err := worker.openSession(ctx, destination, port, firstPayloadTimeout) + return openedSession{stream: connection}, err +} + +func (p *Pool) scheduleIdle(worker *carrierWorker) { + p.mu.Lock() + defer p.mu.Unlock() + if p.closed || !p.containsWorkerLocked(worker) || worker.activeSessions() != 0 { + return + } + p.stopIdleLocked(worker) + p.idle[worker] = p.options.AfterFunc(p.options.IdleTimeout, func() { + p.closeIfIdle(worker) + }) +} + +func (p *Pool) closeIfIdle(worker *carrierWorker) { + p.mu.Lock() + if p.closed || !p.containsWorkerLocked(worker) || worker.activeSessions() != 0 { + p.mu.Unlock() + return + } + p.removeWorkerLocked(worker) + p.mu.Unlock() + worker.close(nil) +} + +func (p *Pool) stopIdleLocked(worker *carrierWorker) { + if timer := p.idle[worker]; timer != nil { + timer.Stop() + delete(p.idle, worker) + } +} + +func (p *Pool) removeWorker(worker *carrierWorker) { + p.mu.Lock() + p.removeWorkerLocked(worker) + p.mu.Unlock() +} + +func (p *Pool) removeWorkerLocked(worker *carrierWorker) { + p.stopIdleLocked(worker) + for index, candidate := range p.workers { + if candidate == worker { + p.workers = append(p.workers[:index], p.workers[index+1:]...) + return + } + } +} + +func (p *Pool) containsWorkerLocked(worker *carrierWorker) bool { + for _, candidate := range p.workers { + if candidate == worker { + return true + } + } + return false +} + +func (p *Pool) activeSessions() int { + p.mu.Lock() + defer p.mu.Unlock() + total := 0 + for _, worker := range p.workers { + total += worker.activeSessions() + } + return total +} + +func (p *Pool) workerCount() int { + p.mu.Lock() + defer p.mu.Unlock() + return len(p.workers) +} + +func (p *Pool) Close() error { + p.closeOnce.Do(func() { + p.mu.Lock() + p.closed = true + close(p.done) + dialCancel := p.dialCancel + workers := append([]*carrierWorker(nil), p.workers...) + p.workers = nil + for worker := range p.idle { + p.stopIdleLocked(worker) + } + p.mu.Unlock() + if dialCancel != nil { + dialCancel() + } + for _, worker := range workers { + worker.close(ErrPoolClosed) + } + }) + return nil +} diff --git a/transport/muxcool/pool_test.go b/transport/muxcool/pool_test.go new file mode 100644 index 0000000000..04b388ae55 --- /dev/null +++ b/transport/muxcool/pool_test.go @@ -0,0 +1,708 @@ +package muxcool + +import ( + "bytes" + "context" + "errors" + "net" + "sync" + "testing" + "time" +) + +type fakeCarrierDialer struct { + mu sync.Mutex + calls int + carriers []*testCarrier + errors []error +} + +type blockingCarrierDialer struct { + mu sync.Mutex + started chan struct{} + startedOnce sync.Once + release chan struct{} + carrier *testCarrier + calls int +} + +func newBlockingCarrierDialer() *blockingCarrierDialer { + return &blockingCarrierDialer{ + started: make(chan struct{}), + release: make(chan struct{}), + carrier: newTestCarrier(), + } +} + +func (d *blockingCarrierDialer) dial(context.Context) (net.Conn, error) { + d.mu.Lock() + d.calls++ + d.mu.Unlock() + d.startedOnce.Do(func() { close(d.started) }) + <-d.release + return d.carrier, nil +} + +func (d *blockingCarrierDialer) callCount() int { + d.mu.Lock() + defer d.mu.Unlock() + return d.calls +} + +func (d *fakeCarrierDialer) dial(context.Context) (net.Conn, error) { + d.mu.Lock() + defer d.mu.Unlock() + d.calls++ + if len(d.errors) > 0 { + err := d.errors[0] + d.errors = d.errors[1:] + if err != nil { + return nil, err + } + } + carrier := newTestCarrier() + d.carriers = append(d.carriers, carrier) + return carrier, nil +} + +func (d *fakeCarrierDialer) callCount() int { + d.mu.Lock() + defer d.mu.Unlock() + return d.calls +} + +func testPoolOptions() Options { + return Options{ + MaxConcurrency: 8, + MaxConnections: 128, + FirstPayloadTimeout: time.Hour, + IdleTimeout: time.Hour, + } +} + +func TestPoolReusesCarrierAndExpandsAtActiveLimit(t *testing.T) { + dialer := &fakeCarrierDialer{} + options := testPoolOptions() + options.MaxConcurrency = 2 + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + first, err := pool.DialContext(context.Background(), "one.example", 80) + if err != nil { + t.Fatal(err) + } + second, err := pool.DialContext(context.Background(), "two.example", 80) + if err != nil { + t.Fatal(err) + } + third, err := pool.DialContext(context.Background(), "three.example", 80) + if err != nil { + t.Fatal(err) + } + defer first.Close() + defer second.Close() + defer third.Close() + if got := dialer.callCount(); got != 2 { + t.Fatalf("carrier dials = %d, want 2", got) + } +} + +func TestPoolMaxCarriersWaitsForReusableCarrier(t *testing.T) { + dialer := &fakeCarrierDialer{} + options := testPoolOptions() + options.MaxConcurrency = 1 + options.CarrierLimiter = NewCarrierLimiter(1) + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + first, err := pool.DialContext(context.Background(), "one.example", 80) + if err != nil { + t.Fatal(err) + } + + result := make(chan net.Conn, 1) + errors := make(chan error, 1) + go func() { + conn, err := pool.DialContext(context.Background(), "two.example", 80) + if err != nil { + errors <- err + return + } + result <- conn + }() + + select { + case conn := <-result: + _ = conn.Close() + t.Fatal("second session opened while the only carrier was full") + case err := <-errors: + t.Fatalf("second session failed while waiting: %v", err) + case <-time.After(50 * time.Millisecond): + } + if got := dialer.callCount(); got != 1 { + t.Fatalf("carrier dials while capped = %d, want 1", got) + } + + if err := first.Close(); err != nil { + t.Fatal(err) + } + select { + case second := <-result: + defer second.Close() + case err := <-errors: + t.Fatalf("second session after capacity release: %v", err) + case <-time.After(time.Second): + t.Fatal("second session did not reuse the carrier after capacity was released") + } + if got := dialer.callCount(); got != 1 { + t.Fatalf("carrier dials after reuse = %d, want 1", got) + } +} + +func TestPoolMaxCarriersWaitHonorsContext(t *testing.T) { + dialer := &fakeCarrierDialer{} + options := testPoolOptions() + options.MaxConcurrency = 1 + options.CarrierLimiter = NewCarrierLimiter(1) + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + first, err := pool.DialContext(context.Background(), "one.example", 80) + if err != nil { + t.Fatal(err) + } + defer first.Close() + + cause := errors.New("carrier wait canceled") + ctx, cancel := context.WithCancelCause(context.Background()) + cancel(cause) + conn, err := pool.DialContext(ctx, "two.example", 80) + if conn != nil || !errors.Is(err, cause) { + t.Fatalf("canceled DialContext = (%v, %v), want (nil, %v)", conn, err, cause) + } + if got := dialer.callCount(); got != 1 { + t.Fatalf("carrier dials after canceled wait = %d, want 1", got) + } +} + +func TestCarrierLimiterIsSharedAcrossPools(t *testing.T) { + dialer := &fakeCarrierDialer{} + limiter := NewCarrierLimiter(1) + options := testPoolOptions() + options.MaxConcurrency = 1 + options.CarrierLimiter = limiter + firstPool := NewPool(dialer.dial, options) + secondPool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = firstPool.Close(); _ = secondPool.Close() }) + + first, err := firstPool.DialContext(context.Background(), "one.example", 80) + if err != nil { + t.Fatal(err) + } + + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + if conn, err := secondPool.DialContext(ctx, "two.example", 80); conn != nil || !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("second pool DialContext = (%v, %v), want deadline exceeded", conn, err) + } + if got := dialer.callCount(); got != 1 { + t.Fatalf("shared carrier dials = %d, want 1", got) + } + + if err := first.Close(); err != nil { + t.Fatal(err) + } + if err := firstPool.Close(); err != nil { + t.Fatal(err) + } + second, err := secondPool.DialContext(context.Background(), "two.example", 80) + if err != nil { + t.Fatal(err) + } + defer second.Close() + if got := dialer.callCount(); got != 2 { + t.Fatalf("carrier dials after release = %d, want 2", got) + } +} + +func TestCarrierLimiterReleasesFailedDialReservation(t *testing.T) { + dialErr := errors.New("dial failed") + dialer := &fakeCarrierDialer{errors: []error{dialErr}} + options := testPoolOptions() + options.CarrierLimiter = NewCarrierLimiter(1) + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + if _, err := pool.DialContext(context.Background(), "failure.example", 80); !errors.Is(err, dialErr) { + t.Fatalf("first dial error = %v", err) + } + conn, err := pool.DialContext(context.Background(), "success.example", 80) + if err != nil { + t.Fatalf("second dial: %v", err) + } + defer conn.Close() + if got := dialer.callCount(); got != 2 { + t.Fatalf("carrier dials = %d, want 2", got) + } +} + +func TestPoolMaxCarriersIsStrictUnderConcurrency(t *testing.T) { + dialer := &fakeCarrierDialer{} + options := testPoolOptions() + options.MaxConcurrency = 1 + options.CarrierLimiter = NewCarrierLimiter(2) + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + const callers = 16 + results := make(chan net.Conn, callers) + errors := make(chan error, callers) + start := make(chan struct{}) + for i := 0; i < callers; i++ { + go func() { + <-start + conn, err := pool.DialContext(context.Background(), "shared.example", 443) + if err != nil { + errors <- err + return + } + results <- conn + }() + } + close(start) + + opened := make([]net.Conn, 0, 2) + for len(opened) < 2 { + select { + case conn := <-results: + opened = append(opened, conn) + case err := <-errors: + t.Fatalf("concurrent dial failed: %v", err) + case <-time.After(time.Second): + t.Fatal("initial carrier capacity did not open") + } + } + select { + case conn := <-results: + _ = conn.Close() + t.Fatal("session exceeded the configured physical carrier capacity") + case err := <-errors: + t.Fatalf("concurrent dial failed while waiting: %v", err) + case <-time.After(50 * time.Millisecond): + } + if got := dialer.callCount(); got != 2 { + t.Fatalf("carrier dials at concurrency limit = %d, want 2", got) + } + + completed := len(opened) + for _, conn := range opened { + if err := conn.Close(); err != nil { + t.Fatal(err) + } + } + for completed < callers { + select { + case conn := <-results: + completed++ + if err := conn.Close(); err != nil { + t.Fatal(err) + } + case err := <-errors: + t.Fatalf("concurrent dial failed: %v", err) + case <-time.After(time.Second): + t.Fatalf("completed sessions = %d, want %d", completed, callers) + } + } + if got := dialer.callCount(); got != 2 { + t.Fatalf("carrier dials after all sessions = %d, want 2", got) + } +} + +func TestPoolCloseUnblocksMaxCarriersWait(t *testing.T) { + dialer := &fakeCarrierDialer{} + options := testPoolOptions() + options.MaxConcurrency = 1 + options.CarrierLimiter = NewCarrierLimiter(1) + pool := NewPool(dialer.dial, options) + + first, err := pool.DialContext(context.Background(), "one.example", 80) + if err != nil { + t.Fatal(err) + } + defer first.Close() + + result := make(chan error, 1) + go func() { + _, err := pool.DialContext(context.Background(), "two.example", 80) + result <- err + }() + select { + case err := <-result: + t.Fatalf("wait returned before pool close: %v", err) + case <-time.After(50 * time.Millisecond): + } + + if err := pool.Close(); err != nil { + t.Fatal(err) + } + select { + case err := <-result: + if !errors.Is(err, ErrPoolClosed) { + t.Fatalf("wait error = %v, want %v", err, ErrPoolClosed) + } + case <-time.After(time.Second): + t.Fatal("Pool.Close did not unblock carrier capacity wait") + } +} + +func TestPoolMaxCarriersWakesAllWaitersForReusableCapacity(t *testing.T) { + dialer := &fakeCarrierDialer{} + options := testPoolOptions() + options.MaxConcurrency = 4 + options.CarrierLimiter = NewCarrierLimiter(1) + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + active := make([]net.PacketConn, 4) + for i := range active { + conn, err := pool.ListenPacketContext(context.Background(), "active.example", 53, [8]byte{}) + if err != nil { + t.Fatal(err) + } + active[i] = conn + } + + const waiters = 3 + results := make(chan net.PacketConn, waiters) + errors := make(chan error, waiters) + for i := 0; i < waiters; i++ { + go func() { + conn, err := pool.ListenPacketContext(context.Background(), "waiting.example", 53, [8]byte{}) + if err != nil { + errors <- err + return + } + results <- conn + }() + } + select { + case conn := <-results: + _ = conn.Close() + t.Fatal("waiter opened while carrier was full") + case err := <-errors: + t.Fatalf("waiter failed: %v", err) + case <-time.After(50 * time.Millisecond): + } + + for i := 0; i < waiters; i++ { + if err := active[i].Close(); err != nil { + t.Fatal(err) + } + } + defer active[3].Close() + for i := 0; i < waiters; i++ { + select { + case conn := <-results: + defer conn.Close() + case err := <-errors: + t.Fatalf("waiter failed after capacity release: %v", err) + case <-time.After(time.Second): + t.Fatalf("woken waiters = %d, want %d", i, waiters) + } + } + if got := dialer.callCount(); got != 1 { + t.Fatalf("carrier dials = %d, want 1", got) + } +} + +func TestPoolSharesCarrierBetweenStreamAndPacketSessions(t *testing.T) { + dialer := &fakeCarrierDialer{} + pool := NewPool(dialer.dial, testPoolOptions()) + t.Cleanup(func() { _ = pool.Close() }) + + stream, err := pool.DialContext(context.Background(), "stream.example", 443) + if err != nil { + t.Fatal(err) + } + packetConn, err := pool.ListenPacketContext(context.Background(), "packet.example", 53, [8]byte{1, 2, 3}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = stream.Close(); _ = packetConn.Close() }) + + if got := dialer.callCount(); got != 1 { + t.Fatalf("carrier dials = %d, want 1", got) + } + if got := pool.activeSessions(); got != 2 { + t.Fatalf("active sessions = %d, want 2", got) + } + if _, err := packetConn.WriteTo([]byte("query"), &net.UDPAddr{IP: net.IPv4(8, 8, 8, 8), Port: 53}); err != nil { + t.Fatal(err) + } + frame, err := DecodeFrame(bytes.NewReader(dialer.carriers[0].bytes())) + if err != nil { + t.Fatal(err) + } + if frame.SessionID != 2 || frame.Status != StatusNew || frame.Network != NetworkUDP || frame.GlobalID != [8]byte{1, 2, 3} { + t.Fatalf("packet frame = %+v", frame) + } +} + +func TestPoolPacketSessionOutlivesDialContext(t *testing.T) { + dialer := &fakeCarrierDialer{} + pool := NewPool(dialer.dial, testPoolOptions()) + t.Cleanup(func() { _ = pool.Close() }) + + ctx, cancel := context.WithCancelCause(context.Background()) + packetConn, err := pool.ListenPacketContext(ctx, "packet.example", 53, [8]byte{1, 2, 3}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = packetConn.Close() }) + cancel(errors.New("dial completed")) + + select { + case <-time.After(20 * time.Millisecond): + } + if got := pool.activeSessions(); got != 1 { + t.Fatalf("active sessions after dial context cancellation = %d, want 1", got) + } + + if _, err := packetConn.WriteTo([]byte("query"), &net.UDPAddr{IP: net.IPv4(8, 8, 8, 8), Port: 53}); err != nil { + t.Fatal(err) + } + frame, err := DecodeFrame(bytes.NewReader(dialer.carriers[0].bytes())) + if err != nil { + t.Fatal(err) + } + if frame.Status != StatusNew || string(frame.Payload) != "query" { + t.Fatalf("packet frame after dial context cancellation = %+v", frame) + } + + dialer.carriers[0].inject(t, Frame{ + SessionID: frame.SessionID, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, + Destination: "8.8.8.8", Port: 53, Payload: []byte("answer"), + }) + buffer := make([]byte, 16) + n, address, err := packetConn.ReadFrom(buffer) + if err != nil { + t.Fatal(err) + } + if string(buffer[:n]) != "answer" || address.String() != "8.8.8.8:53" { + t.Fatalf("packet response = (%q, %v)", buffer[:n], address) + } +} + +func TestPoolRejectsCanceledPacketDialContext(t *testing.T) { + dialer := &fakeCarrierDialer{} + pool := NewPool(dialer.dial, testPoolOptions()) + t.Cleanup(func() { _ = pool.Close() }) + + cause := errors.New("packet dial canceled") + ctx, cancel := context.WithCancelCause(context.Background()) + cancel(cause) + packetConn, err := pool.ListenPacketContext(ctx, "packet.example", 53, [8]byte{}) + if packetConn != nil || !errors.Is(err, cause) { + t.Fatalf("ListenPacketContext = (%v, %v), want (nil, %v)", packetConn, err, cause) + } + if got := pool.activeSessions(); got != 0 { + t.Fatalf("active sessions after canceled dial = %d, want 0", got) + } +} + +func TestPoolRotatesAtLifetimeLimitAndUsesMonotonicIDs(t *testing.T) { + dialer := &fakeCarrierDialer{} + options := testPoolOptions() + options.MaxConnections = 2 + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + for i := 0; i < 2; i++ { + conn, err := pool.DialContext(context.Background(), "echo.example", 443) + if err != nil { + t.Fatal(err) + } + if _, err := conn.Write([]byte{byte(i)}); err != nil { + t.Fatal(err) + } + if err := conn.Close(); err != nil { + t.Fatal(err) + } + waitFor(t, func() bool { return pool.activeSessions() == 0 }) + } + third, err := pool.DialContext(context.Background(), "echo.example", 443) + if err != nil { + t.Fatal(err) + } + defer third.Close() + if got := dialer.callCount(); got != 2 { + t.Fatalf("carrier dials = %d, want 2", got) + } + + reader := bytes.NewReader(dialer.carriers[0].bytes()) + var newIDs []uint16 + for reader.Len() > 0 { + frame, err := DecodeFrame(reader) + if err != nil { + t.Fatalf("decode carrier writes: %v", err) + } + if frame.Status == StatusNew { + newIDs = append(newIDs, frame.SessionID) + } + } + if len(newIDs) != 2 || newIDs[0] != 1 || newIDs[1] != 2 { + t.Fatalf("new session IDs = %v, want [1 2]", newIDs) + } +} + +func TestPoolRemovesIdleCarrierWithFakeClock(t *testing.T) { + dialer := &fakeCarrierDialer{} + clock := &fakeClock{} + options := testPoolOptions() + options.AfterFunc = func(duration time.Duration, fn func()) Timer { + return clock.AfterFunc(duration, fn) + } + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + conn, err := pool.DialContext(context.Background(), "idle.example", 80) + if err != nil { + t.Fatal(err) + } + _ = conn.Close() + waitFor(t, func() bool { return pool.activeSessions() == 0 }) + waitFor(t, func() bool { return clock.timerCount() == 1 }) + clock.FireAll() + waitFor(t, func() bool { return pool.workerCount() == 0 }) + if !dialer.carriers[0].isClosed() { + t.Fatal("idle carrier was not closed") + } +} + +func TestPoolDoesNotRetainFailedCarrierDial(t *testing.T) { + dialErr := errors.New("dial failed") + dialer := &fakeCarrierDialer{errors: []error{dialErr}} + pool := NewPool(dialer.dial, testPoolOptions()) + t.Cleanup(func() { _ = pool.Close() }) + + if _, err := pool.DialContext(context.Background(), "failure.example", 80); !errors.Is(err, dialErr) { + t.Fatalf("first dial error = %v", err) + } + conn, err := pool.DialContext(context.Background(), "success.example", 80) + if err != nil { + t.Fatalf("second dial: %v", err) + } + _ = conn.Close() + if got := dialer.callCount(); got != 2 { + t.Fatalf("carrier dials = %d, want 2", got) + } +} + +func TestPoolCloseDoesNotWaitForCarrierDial(t *testing.T) { + dialer := newBlockingCarrierDialer() + pool := NewPool(dialer.dial, testPoolOptions()) + + dialResult := make(chan error, 1) + go func() { + _, err := pool.DialContext(context.Background(), "blocked.example", 443) + dialResult <- err + }() + + select { + case <-dialer.started: + case <-time.After(time.Second): + t.Fatal("carrier dial did not start") + } + + closeDone := make(chan struct{}) + go func() { + _ = pool.Close() + close(closeDone) + }() + + select { + case <-closeDone: + case <-time.After(100 * time.Millisecond): + close(dialer.release) + <-closeDone + <-dialResult + t.Fatal("Pool.Close blocked on an in-flight carrier dial") + } + + close(dialer.release) + if err := <-dialResult; !errors.Is(err, ErrPoolClosed) { + t.Fatalf("dial error after pool close = %v, want %v", err, ErrPoolClosed) + } + if !dialer.carrier.isClosed() { + t.Fatal("carrier completed after pool close was not closed") + } +} + +func TestPoolCoalescesConcurrentCarrierDials(t *testing.T) { + dialer := newBlockingCarrierDialer() + options := testPoolOptions() + options.MaxConcurrency = 32 + pool := NewPool(dialer.dial, options) + t.Cleanup(func() { _ = pool.Close() }) + + const callers = 16 + start := make(chan struct{}) + results := make(chan net.Conn, callers) + errors := make(chan error, callers) + var ready sync.WaitGroup + ready.Add(callers) + for i := 0; i < callers; i++ { + go func() { + ready.Done() + <-start + conn, err := pool.DialContext(context.Background(), "shared.example", 443) + if err != nil { + errors <- err + return + } + results <- conn + }() + } + ready.Wait() + close(start) + + select { + case <-dialer.started: + case <-time.After(time.Second): + t.Fatal("carrier dial did not start") + } + + select { + case <-time.After(100 * time.Millisecond): + if got := dialer.callCount(); got != 1 { + close(dialer.release) + t.Fatalf("concurrent carrier dials = %d, want 1", got) + } + } + close(dialer.release) + + for i := 0; i < callers; i++ { + select { + case err := <-errors: + t.Fatalf("concurrent dial failed: %v", err) + case conn := <-results: + _ = conn.Close() + case <-time.After(time.Second): + t.Fatal("concurrent dial did not complete") + } + } + if got := dialer.callCount(); got != 1 { + t.Fatalf("carrier dials = %d, want 1", got) + } +} + +func waitFor(t *testing.T, condition func() bool) { + t.Helper() + deadline := time.Now().Add(time.Second) + for !condition() && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + if !condition() { + t.Fatal("condition was not satisfied before timeout") + } +} diff --git a/transport/muxcool/process_e2e_test.go b/transport/muxcool/process_e2e_test.go new file mode 100644 index 0000000000..1851cf8124 --- /dev/null +++ b/transport/muxcool/process_e2e_test.go @@ -0,0 +1,426 @@ +//go:build integration + +package muxcool_test + +import ( + "bytes" + "context" + "fmt" + "io" + "net" + "os" + "os/exec" + "path/filepath" + "runtime" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/metacubex/mihomo/transport/socks5" +) + +const processE2EUUID = "9d0cb9d0-964f-4ef6-897d-6c6b3ccf9e68" + +func TestMuxCoolProcessMaxCarriers(t *testing.T) { + const ( + maxCarriers = 2 + maxConcurrency = 8 + sessionCount = maxCarriers * maxConcurrency + overflowCount = 4 + ) + + binary := buildMihomoProcessE2E(t) + echoAddress, echoAccepts := startProcessE2EEcho(t) + serverPort := reserveProcessE2EPort(t) + relay := startProcessE2ERelay(t, net.JoinHostPort("127.0.0.1", fmt.Sprint(serverPort))) + clientPort := reserveProcessE2EPort(t) + + tempDir := t.TempDir() + serverConfig := filepath.Join(tempDir, "server.yaml") + clientConfig := filepath.Join(tempDir, "client.yaml") + writeProcessE2EFile(t, serverConfig, fmt.Sprintf(` +mode: rule +log-level: debug +listeners: + - name: vless-in + type: vless + listen: 127.0.0.1 + port: %d + allow-insecure: true + users: + - username: process-e2e + uuid: %s +rules: + - MATCH,DIRECT +`, serverPort, processE2EUUID)) + writeProcessE2EFile(t, clientConfig, fmt.Sprintf(` +mixed-port: %d +allow-lan: false +mode: rule +log-level: debug +proxies: + - name: vless-mux-cool + type: vless + server: 127.0.0.1 + port: %d + uuid: %s + network: tcp + udp: true + mux.cool: + enabled: true + max-concurrency: %d + max-connections: 128 + max-carriers: %d +rules: + - MATCH,vless-mux-cool +`, clientPort, relay.port(), processE2EUUID, maxConcurrency, maxCarriers)) + + server := startMihomoProcessE2E(t, binary, serverConfig, filepath.Join(tempDir, "server-home")) + waitProcessE2ETCP(t, server, net.JoinHostPort("127.0.0.1", fmt.Sprint(serverPort))) + client := startMihomoProcessE2E(t, binary, clientConfig, filepath.Join(tempDir, "client-home")) + clientAddress := net.JoinHostPort("127.0.0.1", fmt.Sprint(clientPort)) + waitProcessE2ETCP(t, client, clientAddress) + + connections := openProcessE2EWave(t, clientAddress, echoAddress, 0, sessionCount) + t.Cleanup(func() { closeProcessE2EConnections(connections) }) + if got := relay.total.Load(); got != maxCarriers { + t.Fatalf("physical carrier connections at capacity = %d, want %d", got, maxCarriers) + } + if got := relay.maxActive.Load(); got != maxCarriers { + t.Fatalf("peak active carrier connections = %d, want %d", got, maxCarriers) + } + + overflowResults := startProcessE2EWave(clientAddress, echoAddress, 1, overflowCount) + select { + case result := <-overflowResults: + if result.connection != nil { + _ = result.connection.Close() + } + t.Fatalf("overflow session completed before carrier capacity was released: %v", result.err) + case <-time.After(250 * time.Millisecond): + } + if got := relay.total.Load(); got != maxCarriers { + t.Fatalf("physical carrier connections while capped = %d, want %d", got, maxCarriers) + } + + closeProcessE2EConnections(connections[:overflowCount]) + connections = connections[overflowCount:] + connections = append(connections, collectProcessE2EConnections(t, overflowResults, overflowCount)...) + if got := relay.total.Load(); got != maxCarriers { + t.Fatalf("physical carrier connections after capacity reuse = %d, want %d", got, maxCarriers) + } + closeProcessE2EConnections(connections) + connections = nil + + connections = openProcessE2EWave(t, clientAddress, echoAddress, 2, sessionCount) + closeProcessE2EConnections(connections) + connections = nil + + wantLogical := 2*sessionCount + overflowCount + if got := echoAccepts.Load(); got != int32(wantLogical) { + t.Fatalf("logical target connections = %d, want %d", got, wantLogical) + } + if got := relay.total.Load(); got != maxCarriers { + t.Fatalf("physical carrier connections after reuse = %d, want %d", got, maxCarriers) + } + t.Logf( + "logical sessions=%d physical carriers=%d peak active carriers=%d", + echoAccepts.Load(), relay.total.Load(), relay.maxActive.Load(), + ) +} + +func openProcessE2EWave(t *testing.T, clientAddress, targetAddress string, wave, count int) []net.Conn { + t.Helper() + return collectProcessE2EConnections(t, startProcessE2EWave(clientAddress, targetAddress, wave, count), count) +} + +type processE2EResult struct { + connection net.Conn + err error +} + +func startProcessE2EWave(clientAddress, targetAddress string, wave, count int) <-chan processE2EResult { + results := make(chan processE2EResult, count) + start := make(chan struct{}) + for index := 0; index < count; index++ { + go func(index int) { + <-start + connection, err := net.DialTimeout("tcp", clientAddress, 5*time.Second) + if err != nil { + results <- processE2EResult{err: fmt.Errorf("dial mixed listener: %w", err)} + return + } + if err := connection.SetDeadline(time.Now().Add(10 * time.Second)); err != nil { + _ = connection.Close() + results <- processE2EResult{err: err} + return + } + if _, err := socks5.ClientHandshake(connection, socks5.ParseAddr(targetAddress), socks5.CmdConnect, nil); err != nil { + _ = connection.Close() + results <- processE2EResult{err: fmt.Errorf("SOCKS5 handshake: %w", err)} + return + } + payload := []byte(fmt.Sprintf("wave-%d-session-%d", wave, index)) + if _, err := connection.Write(payload); err != nil { + _ = connection.Close() + results <- processE2EResult{err: fmt.Errorf("write payload: %w", err)} + return + } + response := make([]byte, len(payload)) + if _, err := io.ReadFull(connection, response); err != nil { + _ = connection.Close() + results <- processE2EResult{err: fmt.Errorf("read payload: %w", err)} + return + } + if !bytes.Equal(response, payload) { + _ = connection.Close() + results <- processE2EResult{err: fmt.Errorf("response = %q, want %q", response, payload)} + return + } + _ = connection.SetDeadline(time.Time{}) + results <- processE2EResult{connection: connection} + }(index) + } + close(start) + return results +} + +func collectProcessE2EConnections(t *testing.T, results <-chan processE2EResult, count int) []net.Conn { + t.Helper() + connections := make([]net.Conn, 0, count) + for index := 0; index < count; index++ { + current := <-results + if current.err != nil { + closeProcessE2EConnections(connections) + t.Fatal(current.err) + } + connections = append(connections, current.connection) + } + return connections +} + +func closeProcessE2EConnections(connections []net.Conn) { + for _, connection := range connections { + _ = connection.Close() + } +} + +type processE2ERelay struct { + listener net.Listener + backend string + total atomic.Int32 + active atomic.Int32 + maxActive atomic.Int32 +} + +func startProcessE2ERelay(t *testing.T, backend string) *processE2ERelay { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + relay := &processE2ERelay{listener: listener, backend: backend} + t.Cleanup(func() { _ = listener.Close() }) + go relay.acceptLoop() + return relay +} + +func (r *processE2ERelay) port() int { + return r.listener.Addr().(*net.TCPAddr).Port +} + +func (r *processE2ERelay) acceptLoop() { + for { + client, err := r.listener.Accept() + if err != nil { + return + } + r.total.Add(1) + active := r.active.Add(1) + for { + peak := r.maxActive.Load() + if active <= peak || r.maxActive.CompareAndSwap(peak, active) { + break + } + } + go r.forward(client) + } +} + +func (r *processE2ERelay) forward(client net.Conn) { + defer r.active.Add(-1) + backend, err := net.DialTimeout("tcp", r.backend, 5*time.Second) + if err != nil { + _ = client.Close() + return + } + + done := make(chan struct{}, 2) + go func() { + _, _ = io.Copy(backend, client) + done <- struct{}{} + }() + go func() { + _, _ = io.Copy(client, backend) + done <- struct{}{} + }() + <-done + _ = client.Close() + _ = backend.Close() + <-done +} + +func startProcessE2EEcho(t *testing.T) (string, *atomic.Int32) { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + var accepts atomic.Int32 + t.Cleanup(func() { _ = listener.Close() }) + go func() { + for { + connection, err := listener.Accept() + if err != nil { + return + } + accepts.Add(1) + go func() { + defer connection.Close() + _, _ = io.Copy(connection, connection) + }() + } + }() + return listener.Addr().String(), &accepts +} + +type lockedProcessE2EBuffer struct { + mu sync.Mutex + bytes.Buffer +} + +func (b *lockedProcessE2EBuffer) Write(payload []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.Buffer.Write(payload) +} + +func (b *lockedProcessE2EBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.Buffer.String() +} + +type mihomoProcessE2E struct { + cancel context.CancelFunc + done chan struct{} + mu sync.Mutex + err error + output lockedProcessE2EBuffer +} + +func startMihomoProcessE2E(t *testing.T, binary, config, home string) *mihomoProcessE2E { + t.Helper() + if err := os.MkdirAll(home, 0o755); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + process := &mihomoProcessE2E{cancel: cancel, done: make(chan struct{})} + command := exec.CommandContext(ctx, binary, "-d", home, "-f", config) + command.Stdout = &process.output + command.Stderr = &process.output + if err := command.Start(); err != nil { + cancel() + t.Fatal(err) + } + go func() { + err := command.Wait() + process.mu.Lock() + process.err = err + process.mu.Unlock() + close(process.done) + }() + t.Cleanup(func() { + cancel() + select { + case <-process.done: + case <-time.After(5 * time.Second): + t.Errorf("mihomo process did not stop\n%s", process.output.String()) + } + if t.Failed() { + t.Log(process.output.String()) + } + }) + return process +} + +func (p *mihomoProcessE2E) waitError() error { + p.mu.Lock() + defer p.mu.Unlock() + return p.err +} + +func waitProcessE2ETCP(t *testing.T, process *mihomoProcessE2E, address string) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + ticker := time.NewTicker(25 * time.Millisecond) + defer ticker.Stop() + for { + connection, err := (&net.Dialer{}).DialContext(ctx, "tcp", address) + if err == nil { + _ = connection.Close() + return + } + select { + case <-process.done: + t.Fatalf("mihomo exited before listening on %s: %v\n%s", address, process.waitError(), process.output.String()) + case <-ctx.Done(): + t.Fatalf("mihomo did not listen on %s: %v\n%s", address, context.Cause(ctx), process.output.String()) + case <-ticker.C: + } + } +} + +func buildMihomoProcessE2E(t *testing.T) string { + t.Helper() + if binary := os.Getenv("MIHOMO_E2E_BINARY"); binary != "" { + return binary + } + root, err := filepath.Abs(filepath.Join("..", "..")) + if err != nil { + t.Fatal(err) + } + binary := filepath.Join(t.TempDir(), "mihomo") + if runtime.GOOS == "windows" { + binary += ".exe" + } + command := exec.Command("go", "build", "-trimpath", "-o", binary, ".") + command.Dir = root + output, err := command.CombinedOutput() + if err != nil { + t.Fatalf("build mihomo: %v\n%s", err, output) + } + return binary +} + +func reserveProcessE2EPort(t *testing.T) int { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + port := listener.Addr().(*net.TCPAddr).Port + if err := listener.Close(); err != nil { + t.Fatal(err) + } + return port +} + +func writeProcessE2EFile(t *testing.T, path, content string) { + t.Helper() + if err := os.WriteFile(path, []byte(content), 0o600); err != nil { + t.Fatal(err) + } +} diff --git a/transport/muxcool/server.go b/transport/muxcool/server.go new file mode 100644 index 0000000000..4e23385344 --- /dev/null +++ b/transport/muxcool/server.go @@ -0,0 +1,447 @@ +package muxcool + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "net/netip" + "sync" + "sync/atomic" + "time" + + "github.com/metacubex/sing/common/auth" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" +) + +const ( + defaultXUDPIdleTimeout = time.Minute + defaultMaxSessionsPerCarrier = 1024 +) + +var ( + ErrServerClosed = errors.New("mux.cool server is closed") + ErrDuplicateSessionID = errors.New("mux.cool duplicate session ID") + ErrSessionQueueFull = errors.New("mux.cool session receive queue is full") +) + +type ServerHandler interface { + NewConnection(context.Context, net.Conn, M.Metadata) error + NewPacketConnection(context.Context, N.PacketConn, M.Metadata) error +} + +type ServerTimer interface { + Stop() bool +} + +type ServerOptions struct { + XUDPIdleTimeout time.Duration + MaxSessionsPerCarrier int + AfterFunc func(time.Duration, func()) ServerTimer +} + +type ServerRuntime struct { + options ServerOptions + + mu sync.Mutex + closed bool + carriers map[*serverCarrier]struct{} + flows map[xudpFlowKey]*serverPacketFlow + wg sync.WaitGroup + close sync.Once +} + +func NewServerRuntime(options ServerOptions) *ServerRuntime { + if options.XUDPIdleTimeout <= 0 { + options.XUDPIdleTimeout = defaultXUDPIdleTimeout + } + if options.MaxSessionsPerCarrier <= 0 { + options.MaxSessionsPerCarrier = defaultMaxSessionsPerCarrier + } + if options.AfterFunc == nil { + options.AfterFunc = func(duration time.Duration, callback func()) ServerTimer { + return time.AfterFunc(duration, callback) + } + } + return &ServerRuntime{ + options: options, + carriers: make(map[*serverCarrier]struct{}), + flows: make(map[xudpFlowKey]*serverPacketFlow), + } +} + +func (r *ServerRuntime) Serve(ctx context.Context, conn net.Conn, metadata M.Metadata, handler ServerHandler) error { + carrier := newServerCarrier(r, ctx, conn, metadata, handler) + r.mu.Lock() + if r.closed { + r.mu.Unlock() + _ = conn.Close() + return ErrServerClosed + } + r.carriers[carrier] = struct{}{} + r.wg.Add(1) + r.mu.Unlock() + + defer func() { + r.mu.Lock() + delete(r.carriers, carrier) + r.mu.Unlock() + r.wg.Done() + }() + return carrier.serve() +} + +func (r *ServerRuntime) Close() error { + r.close.Do(func() { + r.mu.Lock() + r.closed = true + carriers := make([]*serverCarrier, 0, len(r.carriers)) + for carrier := range r.carriers { + carriers = append(carriers, carrier) + } + flows := make([]*serverPacketFlow, 0, len(r.flows)) + for _, flow := range r.flows { + flows = append(flows, flow) + } + r.mu.Unlock() + + for _, carrier := range carriers { + carrier.closeWithError(ErrServerClosed) + } + for _, flow := range flows { + flow.closeWithError(ErrServerClosed) + } + r.wg.Wait() + }) + return nil +} + +type xudpFlowKey struct { + principal string + globalID [8]byte +} + +func (r *ServerRuntime) attachXUDP( + ctx context.Context, + carrier *serverCarrier, + frame decodedFrame, +) (*serverPacketAttachment, error) { + key := xudpFlowKey{principal: serverPrincipal(ctx, carrier.metadata), globalID: frame.GlobalID} + destination := socksaddrFromFrame(frame.Frame) + + r.mu.Lock() + if r.closed { + r.mu.Unlock() + frame.releasePayload() + return nil, ErrServerClosed + } + flow := r.flows[key] + if flow == nil { + flowContext, cancel := context.WithCancelCause(context.WithoutCancel(ctx)) + flow = newServerPacketFlow(r, key, flowContext, cancel, carrier.handler, carrier.metadata.Source, destination, true) + r.flows[key] = flow + r.wg.Add(1) + go func() { + defer r.wg.Done() + flow.runHandler() + }() + } else if flow.destination != destination { + r.mu.Unlock() + frame.releasePayload() + return nil, fmt.Errorf("mux.cool XUDP GlobalID target mismatch: have %s, got %s", flow.destination, destination) + } + r.mu.Unlock() + + attachment, err := flow.attach(carrier, frame.SessionID) + if err != nil { + frame.releasePayload() + return nil, err + } + return attachment, nil +} + +func (r *ServerRuntime) removeFlow(key xudpFlowKey, flow *serverPacketFlow) { + r.mu.Lock() + if r.flows[key] == flow { + delete(r.flows, key) + } + r.mu.Unlock() +} + +func serverPrincipal(ctx context.Context, metadata M.Metadata) string { + if user, loaded := auth.UserFromContext[string](ctx); loaded { + return "user:" + user + } + return "source:" + metadata.Source.String() +} + +type serverSession interface { + deliverDecodedFrame(decodedFrame) error + finish(error, bool) +} + +type serverCarrier struct { + runtime *ServerRuntime + ctx context.Context + cancel context.CancelCauseFunc + conn net.Conn + metadata M.Metadata + handler ServerHandler + + writeMu sync.Mutex + writeBuffer []byte + mu sync.Mutex + sessions map[uint16]serverSession + closed bool + closedFast atomic.Bool + closeErr error + closeOnce sync.Once +} + +func newServerCarrier(runtime *ServerRuntime, parent context.Context, conn net.Conn, metadata M.Metadata, handler ServerHandler) *serverCarrier { + ctx, cancel := context.WithCancelCause(parent) + return &serverCarrier{ + runtime: runtime, + ctx: ctx, + cancel: cancel, + conn: conn, + metadata: metadata, + handler: handler, + sessions: make(map[uint16]serverSession), + } +} + +func (c *serverCarrier) serve() error { + ctxDone := c.ctx.Done() + if ctxDone != nil { + go func() { + <-ctxDone + c.closeWithError(context.Cause(c.ctx)) + }() + } + + metadataBuffer := make([]byte, MaxMetadataSize) + for { + frame, err := decodeFramePooled(c.conn, metadataBuffer) + if err != nil { + c.closeWithError(err) + return err + } + if err := c.handleFrame(frame); err != nil { + c.closeWithError(err) + return err + } + } +} + +func (c *serverCarrier) handleFrame(frame decodedFrame) error { + switch frame.Status { + case StatusKeepAlive: + frame.releasePayload() + return nil + case StatusNew: + return c.handleNew(frame) + case StatusKeep, StatusEnd: + c.mu.Lock() + session := c.sessions[frame.SessionID] + c.mu.Unlock() + if session == nil { + frame.releasePayload() + if frame.Status != StatusEnd { + return c.writeFrame(Frame{SessionID: frame.SessionID, Status: StatusEnd, Option: OptionError}) + } + return nil + } + if frame.Status == StatusEnd { + frame.releasePayload() + session.finish(remoteFrameError(frame.Frame), false) + return nil + } + if err := session.deliverDecodedFrame(frame); err != nil { + session.finish(err, true) + } + return nil + default: + frame.releasePayload() + return protocolError("server", fmt.Errorf("invalid status %d", frame.Status)) + } +} + +func (c *serverCarrier) handleNew(frame decodedFrame) error { + c.mu.Lock() + if c.closed { + c.mu.Unlock() + frame.releasePayload() + return net.ErrClosed + } + if _, exists := c.sessions[frame.SessionID]; exists { + c.mu.Unlock() + frame.releasePayload() + return c.writeFrame(Frame{SessionID: frame.SessionID, Status: StatusEnd, Option: OptionError}) + } + if len(c.sessions) >= c.runtime.options.MaxSessionsPerCarrier { + c.mu.Unlock() + frame.releasePayload() + return c.writeFrame(Frame{SessionID: frame.SessionID, Status: StatusEnd, Option: OptionError}) + } + // Reserve the peer-controlled ID before any handler or flow work. + c.sessions[frame.SessionID] = nil + c.mu.Unlock() + + var ( + session serverSession + err error + ) + switch frame.Network { + case NetworkTCP: + session = newServerStream(c, frame.SessionID, socksaddrFromFrame(frame.Frame)) + if !c.publishReserved(frame.SessionID, session) { + frame.releasePayload() + err = net.ErrClosed + } else if err = session.deliverDecodedFrame(frame); err == nil { + session.(*serverStream).start() + } + case NetworkUDP: + if frame.GlobalID != [8]byte{} { + var attachment *serverPacketAttachment + attachment, err = c.runtime.attachXUDP(c.ctx, c, frame) + if err == nil { + session = attachment + if !c.publishReserved(frame.SessionID, session) { + frame.releasePayload() + err = net.ErrClosed + } else { + err = session.deliverDecodedFrame(frame) + } + } + } else { + flowContext, cancel := context.WithCancelCause(c.ctx) + flow := newServerPacketFlow(nil, xudpFlowKey{}, flowContext, cancel, c.handler, c.metadata.Source, socksaddrFromFrame(frame.Frame), false) + session, err = flow.attach(c, frame.SessionID) + if err == nil { + if !c.publishReserved(frame.SessionID, session) { + frame.releasePayload() + err = net.ErrClosed + } else if err = session.deliverDecodedFrame(frame); err == nil { + go flow.runHandler() + } + } + } + default: + frame.releasePayload() + err = protocolError("server new", fmt.Errorf("invalid network %d", frame.Network)) + } + if err != nil { + c.removeReserved(frame.SessionID) + if session != nil { + session.finish(err, true) + } else { + _ = c.writeFrame(Frame{SessionID: frame.SessionID, Status: StatusEnd, Option: OptionError}) + } + } + return nil +} + +func (c *serverCarrier) publishReserved(id uint16, session serverSession) bool { + c.mu.Lock() + defer c.mu.Unlock() + if reserved, exists := c.sessions[id]; exists && reserved == nil && !c.closed { + c.sessions[id] = session + return true + } + return false +} + +func (c *serverCarrier) removeReserved(id uint16) { + c.mu.Lock() + if session, exists := c.sessions[id]; exists && session == nil { + delete(c.sessions, id) + } + c.mu.Unlock() +} + +func (c *serverCarrier) removeSession(id uint16, expected serverSession) { + c.mu.Lock() + if c.sessions[id] == expected { + delete(c.sessions, id) + } + c.mu.Unlock() +} + +func (c *serverCarrier) writeFrame(frame Frame) error { + c.writeMu.Lock() + encoded, err := encodeFrame(c.writeBuffer, frame) + if err == nil { + c.writeBuffer = encoded[:0] + if c.closedFast.Load() { + c.mu.Lock() + closeErr := c.closeErr + c.mu.Unlock() + if closeErr == nil { + closeErr = net.ErrClosed + } + err = closeErr + } else { + err = writeFull(c.conn, encoded) + } + } + c.writeMu.Unlock() + if err != nil { + // Session close may write an End frame while carrier close is already + // fanning out. Defer the error close so sync.Once is never re-entered. + go c.closeWithError(err) + } + return err +} + +func (c *serverCarrier) closeWithError(cause error) { + c.closeOnce.Do(func() { + if cause == nil { + cause = net.ErrClosed + } + c.mu.Lock() + c.closed = true + c.closeErr = cause + c.closedFast.Store(true) + sessions := make([]serverSession, 0, len(c.sessions)) + for _, session := range c.sessions { + if session != nil { + sessions = append(sessions, session) + } + } + c.sessions = nil + c.mu.Unlock() + + c.cancel(cause) + _ = c.conn.Close() + for _, session := range sessions { + session.finish(cause, false) + } + }) +} + +func socksaddrFromFrame(frame Frame) M.Socksaddr { + if frame.DestinationIP.IsValid() { + return M.Socksaddr{Addr: frame.DestinationIP.Unmap(), Port: frame.Port} + } + if address, err := netip.ParseAddr(frame.Destination); err == nil && address.Zone() == "" { + return M.Socksaddr{Addr: address.Unmap(), Port: frame.Port} + } + return M.Socksaddr{Fqdn: frame.Destination, Port: frame.Port} +} + +func frameTarget(destination M.Socksaddr) (string, netip.Addr, uint16) { + if destination.Addr.IsValid() { + return "", destination.Addr.Unmap(), destination.Port + } + return destination.Fqdn, netip.Addr{}, destination.Port +} + +func remoteFrameError(frame Frame) error { + if frame.Option&OptionError != 0 { + return errors.New("mux.cool remote session error") + } + return io.EOF +} diff --git a/transport/muxcool/server_packet.go b/transport/muxcool/server_packet.go new file mode 100644 index 0000000000..97c528d797 --- /dev/null +++ b/transport/muxcool/server_packet.go @@ -0,0 +1,398 @@ +package muxcool + +import ( + "context" + "errors" + "net" + "os" + "sync" + "sync/atomic" + "time" + + "github.com/metacubex/mihomo/common/net/deadline" + "github.com/metacubex/sing/common/buf" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" +) + +const serverPacketQueueSize = 32 + +type serverPacketMessage struct { + payload []byte + payloadPooled bool + destination M.Socksaddr +} + +func (m *serverPacketMessage) release() { + releasePooledPayload(m.payload, m.payloadPooled) + m.payload = nil + m.payloadPooled = false +} + +type serverPacketFlow struct { + runtime *ServerRuntime + key xudpFlowKey + ctx context.Context + cancel context.CancelCauseFunc + handler ServerHandler + source M.Socksaddr + destination M.Socksaddr + reusable bool + + input chan serverPacketMessage + done chan struct{} + readDeadline deadline.PipeDeadline + writeDeadline deadline.PipeDeadline + writeDeadlineSet atomic.Bool + readOptions N.ReadWaitOptions + + mu sync.Mutex + current *serverPacketAttachment + currentFast atomic.Pointer[serverPacketAttachment] + generation uint64 + idleTimer ServerTimer + closed bool + closedFast atomic.Bool + closeErr error +} + +func newServerPacketFlow( + runtime *ServerRuntime, + key xudpFlowKey, + ctx context.Context, + cancel context.CancelCauseFunc, + handler ServerHandler, + source M.Socksaddr, + destination M.Socksaddr, + reusable bool, +) *serverPacketFlow { + return &serverPacketFlow{ + runtime: runtime, + key: key, + ctx: ctx, + cancel: cancel, + handler: handler, + source: source, + destination: destination, + reusable: reusable, + input: make(chan serverPacketMessage, serverPacketQueueSize), + done: make(chan struct{}), + readDeadline: deadline.MakePipeDeadline(), + writeDeadline: deadline.MakePipeDeadline(), + } +} + +func (f *serverPacketFlow) runHandler() { + err := f.handler.NewPacketConnection(f.ctx, f, M.Metadata{Source: f.source, Destination: f.destination}) + f.closeWithError(err) +} + +func (f *serverPacketFlow) attach(carrier *serverCarrier, id uint16) (*serverPacketAttachment, error) { + f.mu.Lock() + if f.closed { + err := f.closeErr + if err == nil { + err = net.ErrClosed + } + f.mu.Unlock() + return nil, err + } + if f.idleTimer != nil { + f.idleTimer.Stop() + f.idleTimer = nil + } + f.generation++ + attachment := &serverPacketAttachment{ + flow: f, + carrier: carrier, + id: id, + generation: f.generation, + } + previous := f.current + f.current = attachment + f.currentFast.Store(attachment) + f.mu.Unlock() + + if previous != nil { + previous.finish(nil, true) + } + return attachment, nil +} + +func (f *serverPacketFlow) detach(attachment *serverPacketAttachment, cause error) { + f.mu.Lock() + if f.current != attachment { + f.mu.Unlock() + return + } + f.current = nil + f.currentFast.Store(nil) + if f.closed { + f.mu.Unlock() + return + } + if !f.reusable { + f.mu.Unlock() + f.closeWithError(cause) + return + } + timeout := f.runtime.options.XUDPIdleTimeout + generation := f.generation + f.idleTimer = f.runtime.options.AfterFunc(timeout, func() { + f.expire(generation) + }) + f.mu.Unlock() +} + +func (f *serverPacketFlow) enqueue(attachment *serverPacketAttachment, frame decodedFrame) error { + if f.currentFast.Load() != attachment { + frame.releasePayload() + return net.ErrClosed + } + destination := socksaddrFromFrame(frame.Frame) + if !destination.IsValid() { + destination = f.destination + } + message := serverPacketMessage{ + payload: frame.Payload, + payloadPooled: frame.payloadPooled, + destination: destination, + } + select { + case f.input <- message: + return nil + case <-f.done: + message.release() + return f.terminalError() + default: + message.release() + return ErrSessionQueueFull + } +} + +func (f *serverPacketFlow) nextMessage() (serverPacketMessage, error) { + select { + case message := <-f.input: + return message, nil + default: + } + select { + case message := <-f.input: + return message, nil + case <-f.done: + return serverPacketMessage{}, f.terminalError() + case <-f.readDeadline.Wait(): + return serverPacketMessage{}, os.ErrDeadlineExceeded + } +} + +func (f *serverPacketFlow) ReadPacket(buffer *buf.Buffer) (M.Socksaddr, error) { + message, err := f.nextMessage() + if err != nil { + return M.Socksaddr{}, err + } + defer message.release() + if _, err := buffer.Write(message.payload); err != nil { + return M.Socksaddr{}, err + } + return message.destination, nil +} + +func (f *serverPacketFlow) InitializeReadWaiter(options N.ReadWaitOptions) bool { + f.readOptions = options + return false +} + +func (f *serverPacketFlow) WaitReadPacket() (*buf.Buffer, M.Socksaddr, error) { + message, err := f.nextMessage() + if err != nil { + return nil, M.Socksaddr{}, err + } + defer message.release() + buffer := f.readOptions.NewPacketBuffer() + if _, err := buffer.Write(message.payload); err != nil { + buffer.Release() + return nil, M.Socksaddr{}, err + } + f.readOptions.PostReturn(buffer) + return buffer, message.destination, nil +} + +func (f *serverPacketFlow) WritePacket(buffer *buf.Buffer, destination M.Socksaddr) error { + if f.closedFast.Load() { + return f.terminalError() + } + if f.writeDeadlineSet.Load() { + select { + case <-f.writeDeadline.Wait(): + return os.ErrDeadlineExceeded + default: + } + } + attachment := f.currentFast.Load() + if attachment == nil { + // A reusable XUDP backend may emit a late response while detached. The + // packet belongs to no active carrier generation and must be dropped. + return nil + } + err := attachment.writePacket(buffer.Bytes(), destination) + if err == nil { + return nil + } + // A write admitted by the old generation may finish after a rebind. Its + // failure owns only that retired attachment and must not kill the current + // XUDP backend or attachment. + if f.currentFast.Load() != attachment { + return nil + } + return err +} + +func (f *serverPacketFlow) LocalAddr() net.Addr { return muxAddr("mux.cool-udp-server") } + +func (f *serverPacketFlow) Close() error { + f.closeWithError(nil) + return nil +} + +func (f *serverPacketFlow) SetDeadline(value time.Time) error { + f.readDeadline.Set(value) + return f.SetWriteDeadline(value) +} + +func (f *serverPacketFlow) SetReadDeadline(value time.Time) error { + f.readDeadline.Set(value) + return nil +} + +func (f *serverPacketFlow) SetWriteDeadline(value time.Time) error { + if !value.IsZero() { + f.writeDeadlineSet.Store(true) + } + f.writeDeadline.Set(value) + if value.IsZero() { + f.writeDeadlineSet.Store(false) + } + return nil +} + +func (f *serverPacketFlow) closeWithError(cause error) { + current, cause, closed := f.markClosed(cause, 0, false) + if !closed { + return + } + f.finishClose(current, cause) +} + +func (f *serverPacketFlow) expire(generation uint64) { + current, cause, closed := f.markClosed(context.DeadlineExceeded, generation, true) + if !closed { + return + } + f.finishClose(current, cause) +} + +func (f *serverPacketFlow) markClosed(cause error, generation uint64, requireIdleGeneration bool) (*serverPacketAttachment, error, bool) { + if cause == nil { + cause = net.ErrClosed + } + f.mu.Lock() + if f.closed || (requireIdleGeneration && (f.current != nil || f.generation != generation)) { + f.mu.Unlock() + return nil, cause, false + } + f.closed = true + f.closeErr = cause + f.closedFast.Store(true) + if f.idleTimer != nil { + f.idleTimer.Stop() + f.idleTimer = nil + } + current := f.current + f.current = nil + f.currentFast.Store(nil) + f.mu.Unlock() + return current, cause, true +} + +func (f *serverPacketFlow) finishClose(current *serverPacketAttachment, cause error) { + f.cancel(cause) + close(f.done) + if current != nil { + current.finish(cause, true) + } + f.releaseQueued() + if f.runtime != nil { + f.runtime.removeFlow(f.key, f) + } +} + +func (f *serverPacketFlow) terminalError() error { + f.mu.Lock() + defer f.mu.Unlock() + if f.closeErr != nil { + return f.closeErr + } + return net.ErrClosed +} + +func (f *serverPacketFlow) releaseQueued() { + for { + select { + case message := <-f.input: + message.release() + default: + return + } + } +} + +type serverPacketAttachment struct { + flow *serverPacketFlow + carrier *serverCarrier + id uint16 + generation uint64 + closeOnce sync.Once +} + +func (a *serverPacketAttachment) deliverDecodedFrame(frame decodedFrame) error { + if frame.Option&OptionData == 0 || len(frame.Payload) == 0 { + frame.releasePayload() + return nil + } + return a.flow.enqueue(a, frame) +} + +func (a *serverPacketAttachment) writePacket(payload []byte, destination M.Socksaddr) error { + host, ip, port := frameTarget(destination) + return a.carrier.writeFrame(Frame{ + SessionID: a.id, + Status: StatusKeep, + Option: OptionData, + Network: NetworkUDP, + Destination: host, + DestinationIP: ip, + Port: port, + Payload: payload, + }) +} + +func (a *serverPacketAttachment) finish(cause error, sendEnd bool) { + a.closeOnce.Do(func() { + a.carrier.removeSession(a.id, a) + a.flow.detach(a, cause) + if sendEnd { + option := Option(0) + if cause != nil && !errors.Is(cause, net.ErrClosed) && !errors.Is(cause, context.Canceled) { + option = OptionError + } + _ = a.carrier.writeFrame(Frame{SessionID: a.id, Status: StatusEnd, Option: option}) + } + }) +} + +var ( + _ N.PacketConn = (*serverPacketFlow)(nil) + _ N.PacketReadWaiter = (*serverPacketFlow)(nil) + _ serverSession = (*serverPacketAttachment)(nil) +) diff --git a/transport/muxcool/server_stream.go b/transport/muxcool/server_stream.go new file mode 100644 index 0000000000..9c04daff6f --- /dev/null +++ b/transport/muxcool/server_stream.go @@ -0,0 +1,146 @@ +package muxcool + +import ( + "errors" + "io" + "net" + "sync" + + M "github.com/metacubex/sing/common/metadata" +) + +const serverStreamQueueSize = 16 + +type serverStreamMessage struct { + payload []byte + payloadPooled bool +} + +func (m *serverStreamMessage) release() { + releasePooledPayload(m.payload, m.payloadPooled) + m.payload = nil + m.payloadPooled = false +} + +type serverStream struct { + carrier *serverCarrier + id uint16 + destination M.Socksaddr + client net.Conn + peer net.Conn + input chan serverStreamMessage + done chan struct{} + closeOnce sync.Once +} + +func newServerStream(carrier *serverCarrier, id uint16, destination M.Socksaddr) *serverStream { + client, peer := net.Pipe() + return &serverStream{ + carrier: carrier, + id: id, + destination: destination, + client: client, + peer: peer, + input: make(chan serverStreamMessage, serverStreamQueueSize), + done: make(chan struct{}), + } +} + +func (s *serverStream) start() { + go s.writeInput() + go s.readOutput() + go func() { + metadata := M.Metadata{Source: s.carrier.metadata.Source, Destination: s.destination} + if err := s.carrier.handler.NewConnection(s.carrier.ctx, s.client, metadata); err != nil { + s.finish(err, true) + } + }() +} + +func (s *serverStream) deliverDecodedFrame(frame decodedFrame) error { + if frame.Option&OptionData == 0 || len(frame.Payload) == 0 { + frame.releasePayload() + return nil + } + message := serverStreamMessage{payload: frame.Payload, payloadPooled: frame.payloadPooled} + select { + case s.input <- message: + return nil + case <-s.done: + message.release() + return net.ErrClosed + default: + message.release() + return ErrSessionQueueFull + } +} + +func (s *serverStream) writeInput() { + for { + select { + case message := <-s.input: + err := writeFull(s.peer, message.payload) + message.release() + if err != nil { + s.finish(err, true) + return + } + case <-s.done: + s.releaseQueued() + return + } + } +} + +func (s *serverStream) readOutput() { + buffer := sessionBufferPool.Get().(*[MaxPayloadSize]byte) + defer sessionBufferPool.Put(buffer) + for { + n, err := s.peer.Read(buffer[:]) + if n > 0 { + writeErr := s.carrier.writeFrame(Frame{ + SessionID: s.id, + Status: StatusKeep, + Option: OptionData, + Payload: buffer[:n], + }) + if writeErr != nil { + s.finish(writeErr, false) + return + } + } + if err != nil { + s.finish(nil, errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed)) + return + } + } +} + +func (s *serverStream) finish(cause error, sendEnd bool) { + s.closeOnce.Do(func() { + close(s.done) + _ = s.peer.Close() + _ = s.client.Close() + s.carrier.removeSession(s.id, s) + if sendEnd { + option := Option(0) + if cause != nil && !errors.Is(cause, io.EOF) && !errors.Is(cause, net.ErrClosed) { + option = OptionError + } + _ = s.carrier.writeFrame(Frame{SessionID: s.id, Status: StatusEnd, Option: option}) + } + }) +} + +func (s *serverStream) releaseQueued() { + for { + select { + case message := <-s.input: + message.release() + default: + return + } + } +} + +var _ serverSession = (*serverStream)(nil) diff --git a/transport/muxcool/server_test.go b/transport/muxcool/server_test.go new file mode 100644 index 0000000000..22a9f3e9ce --- /dev/null +++ b/transport/muxcool/server_test.go @@ -0,0 +1,606 @@ +package muxcool + +import ( + "context" + "errors" + "io" + "net" + "net/netip" + "os" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/metacubex/sing/common/auth" + "github.com/metacubex/sing/common/buf" + M "github.com/metacubex/sing/common/metadata" + N "github.com/metacubex/sing/common/network" +) + +type testServerHandler struct { + tcp func(context.Context, net.Conn, M.Metadata) error + udp func(context.Context, N.PacketConn, M.Metadata) error +} + +func TestServerPacketFlowWriteDeadlineFastPath(t *testing.T) { + ctx, cancel := context.WithCancelCause(context.Background()) + flow := newServerPacketFlow( + nil, xudpFlowKey{}, ctx, cancel, testServerHandler{}, M.Socksaddr{}, + M.Socksaddr{Fqdn: "deadline.example", Port: 53}, false, + ) + packet := buf.NewSize(1) + defer packet.Release() + _, _ = packet.Write([]byte{1}) + + if err := flow.SetWriteDeadline(time.Now().Add(-time.Second)); err != nil { + t.Fatal(err) + } + if err := flow.WritePacket(packet, flow.destination); !errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatalf("expired write = %v, want deadline exceeded", err) + } + if err := flow.SetWriteDeadline(time.Time{}); err != nil { + t.Fatal(err) + } + if err := flow.WritePacket(packet, flow.destination); err != nil { + t.Fatalf("reset write = %v", err) + } + + flow.closeWithError(io.EOF) + if err := flow.WritePacket(packet, flow.destination); !errors.Is(err, io.EOF) { + t.Fatalf("closed write = %v, want EOF", err) + } +} + +type manualServerTimer struct { + stopped atomic.Bool + callback func() +} + +type blockingErrorConn struct { + started chan struct{} + release chan struct{} + once sync.Once +} + +func (c *blockingErrorConn) Read([]byte) (int, error) { return 0, io.EOF } +func (c *blockingErrorConn) Write([]byte) (int, error) { + c.once.Do(func() { close(c.started) }) + <-c.release + return 0, io.ErrClosedPipe +} +func (c *blockingErrorConn) Close() error { return nil } +func (c *blockingErrorConn) LocalAddr() net.Addr { return muxAddr("local") } +func (c *blockingErrorConn) RemoteAddr() net.Addr { return muxAddr("remote") } +func (c *blockingErrorConn) SetDeadline(time.Time) error { return nil } +func (c *blockingErrorConn) SetReadDeadline(time.Time) error { return nil } +func (c *blockingErrorConn) SetWriteDeadline(time.Time) error { return nil } + +func (t *manualServerTimer) Stop() bool { + return !t.stopped.Swap(true) +} + +func (t *manualServerTimer) Fire() { + t.callback() +} + +func (h testServerHandler) NewConnection(ctx context.Context, conn net.Conn, metadata M.Metadata) error { + return h.tcp(ctx, conn, metadata) +} + +func (h testServerHandler) NewPacketConnection(ctx context.Context, conn N.PacketConn, metadata M.Metadata) error { + return h.udp(ctx, conn, metadata) +} + +func testCarrierMetadata() M.Metadata { + return M.Metadata{Source: M.Socksaddr{Addr: netip.MustParseAddr("192.0.2.10"), Port: 32000}} +} + +func serveTestCarrier(t *testing.T, runtime *ServerRuntime, handler ServerHandler) net.Conn { + return serveTestCarrierContext(t, context.Background(), runtime, handler) +} + +func serveTestCarrierContext(t *testing.T, ctx context.Context, runtime *ServerRuntime, handler ServerHandler) net.Conn { + t.Helper() + client, server := net.Pipe() + done := make(chan error, 1) + go func() { + done <- runtime.Serve(ctx, server, testCarrierMetadata(), handler) + }() + t.Cleanup(func() { + _ = client.Close() + select { + case err := <-done: + if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, net.ErrClosed) { + t.Errorf("Serve: %v", err) + } + case <-time.After(time.Second): + t.Error("Serve did not stop") + } + }) + return client +} + +func writeTestFrame(t *testing.T, conn net.Conn, frame Frame) { + t.Helper() + raw, err := EncodeFrame(frame) + if err != nil { + t.Fatal(err) + } + if err := writeFull(conn, raw); err != nil { + t.Fatal(err) + } +} + +func readTestFrame(t *testing.T, conn net.Conn) Frame { + t.Helper() + if err := conn.SetReadDeadline(time.Now().Add(time.Second)); err != nil { + t.Fatal(err) + } + frame, err := DecodeFrame(conn) + if err != nil { + t.Fatal(err) + } + if err := conn.SetReadDeadline(time.Time{}); err != nil { + t.Fatal(err) + } + return frame +} + +func TestServerRuntimeMultiplexesTCP(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{XUDPIdleTimeout: time.Hour}) + t.Cleanup(func() { _ = runtime.Close() }) + + handler := testServerHandler{ + tcp: func(_ context.Context, conn net.Conn, metadata M.Metadata) error { + if metadata.Destination.Fqdn != "echo.example" || metadata.Destination.Port != 443 { + return errors.New("unexpected destination") + } + payload := make([]byte, 32) + n, err := conn.Read(payload) + if err != nil { + return err + } + _, err = conn.Write(payload[:n]) + return err + }, + udp: func(context.Context, N.PacketConn, M.Metadata) error { return errors.New("unexpected UDP") }, + } + carrier := serveTestCarrier(t, runtime, handler) + + writeTestFrame(t, carrier, Frame{ + SessionID: 7, Status: StatusNew, Option: OptionData, Network: NetworkTCP, + Destination: "echo.example", Port: 443, Payload: []byte("hello"), + }) + response := readTestFrame(t, carrier) + if response.SessionID != 7 || response.Status != StatusKeep || string(response.Payload) != "hello" { + t.Fatalf("response = %+v", response) + } +} + +func TestServerRuntimeRejectsDuplicateSessionIDBeforeDispatch(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{XUDPIdleTimeout: time.Hour}) + t.Cleanup(func() { _ = runtime.Close() }) + var calls atomic.Int32 + release := make(chan struct{}) + started := make(chan struct{}, 2) + handler := testServerHandler{ + tcp: func(_ context.Context, conn net.Conn, _ M.Metadata) error { + calls.Add(1) + started <- struct{}{} + <-release + return conn.Close() + }, + udp: func(context.Context, N.PacketConn, M.Metadata) error { return errors.New("unexpected UDP") }, + } + carrier := serveTestCarrier(t, runtime, handler) + frame := Frame{SessionID: 11, Status: StatusNew, Network: NetworkTCP, Destination: "one.example", Port: 80} + writeTestFrame(t, carrier, frame) + <-started + writeTestFrame(t, carrier, frame) + + response := readTestFrame(t, carrier) + if response.SessionID != 11 || response.Status != StatusEnd || response.Option&OptionError == 0 { + t.Fatalf("duplicate response = %+v", response) + } + if got := calls.Load(); got != 1 { + t.Fatalf("dispatch calls = %d, want 1", got) + } + close(release) +} + +func TestServerRuntimeEnforcesPerCarrierSessionLimit(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{MaxSessionsPerCarrier: 1, XUDPIdleTimeout: time.Hour}) + t.Cleanup(func() { _ = runtime.Close() }) + var calls atomic.Int32 + release := make(chan struct{}) + started := make(chan struct{}, 2) + handler := testServerHandler{ + tcp: func(_ context.Context, conn net.Conn, _ M.Metadata) error { + calls.Add(1) + started <- struct{}{} + <-release + return conn.Close() + }, + udp: func(context.Context, N.PacketConn, M.Metadata) error { return errors.New("unexpected UDP") }, + } + carrier := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, carrier, Frame{SessionID: 1, Status: StatusNew, Network: NetworkTCP, Destination: "one.example", Port: 80}) + <-started + writeTestFrame(t, carrier, Frame{SessionID: 2, Status: StatusNew, Network: NetworkTCP, Destination: "two.example", Port: 80}) + response := readTestFrame(t, carrier) + if response.SessionID != 2 || response.Status != StatusEnd || response.Option&OptionError == 0 { + t.Fatalf("limit response = %+v", response) + } + if got := calls.Load(); got != 1 { + t.Fatalf("dispatch calls = %d, want 1", got) + } + close(release) +} + +func TestServerRuntimePreservesUDPPacketsAndTargets(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{XUDPIdleTimeout: time.Hour}) + t.Cleanup(func() { _ = runtime.Close() }) + handler := testServerHandler{ + tcp: func(context.Context, net.Conn, M.Metadata) error { return errors.New("unexpected TCP") }, + udp: func(_ context.Context, conn N.PacketConn, metadata M.Metadata) error { + if metadata.Destination.Fqdn != "dns.example" || metadata.Destination.Port != 53 { + return errors.New("unexpected destination") + } + packet := buf.NewSize(MaxPayloadSize) + defer packet.Release() + destination, err := conn.ReadPacket(packet) + if err != nil { + return err + } + return conn.WritePacket(packet, destination) + }, + } + carrier := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, carrier, Frame{ + SessionID: 3, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "dns.example", Port: 53, Payload: []byte("query"), + }) + response := readTestFrame(t, carrier) + if response.SessionID != 3 || response.Status != StatusKeep || response.Network != NetworkUDP || + response.Destination != "dns.example" || response.Port != 53 || string(response.Payload) != "query" { + t.Fatalf("response = %+v", response) + } +} + +func TestServerRuntimeRebindsXUDPFlowAcrossCarriers(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{XUDPIdleTimeout: time.Hour}) + t.Cleanup(func() { _ = runtime.Close() }) + var calls atomic.Int32 + handler := testServerHandler{ + tcp: func(context.Context, net.Conn, M.Metadata) error { return errors.New("unexpected TCP") }, + udp: func(_ context.Context, conn N.PacketConn, _ M.Metadata) error { + calls.Add(1) + for index := 0; index < 2; index++ { + packet := buf.NewSize(MaxPayloadSize) + destination, err := conn.ReadPacket(packet) + if err != nil { + packet.Release() + return err + } + err = conn.WritePacket(packet, destination) + packet.Release() + if err != nil { + return err + } + } + return nil + }, + } + globalID := [8]byte{1, 2, 3, 4, 5, 6, 7, 8} + + firstCarrier := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, firstCarrier, Frame{ + SessionID: 1, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "xudp.example", Port: 443, GlobalID: globalID, Payload: []byte("first"), + }) + if response := readTestFrame(t, firstCarrier); string(response.Payload) != "first" { + t.Fatalf("first response = %+v", response) + } + writeTestFrame(t, firstCarrier, Frame{SessionID: 1, Status: StatusEnd}) + + secondCarrier := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, secondCarrier, Frame{ + SessionID: 9, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "xudp.example", Port: 443, GlobalID: globalID, Payload: []byte("second"), + }) + if response := readTestFrame(t, secondCarrier); response.SessionID != 9 || string(response.Payload) != "second" { + t.Fatalf("second response = %+v", response) + } + if got := calls.Load(); got != 1 { + t.Fatalf("XUDP dispatch calls = %d, want 1", got) + } +} + +func TestServerRuntimeXUDPStaleAttachmentCannotDetachCurrent(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{XUDPIdleTimeout: time.Hour}) + t.Cleanup(func() { _ = runtime.Close() }) + var calls atomic.Int32 + handler := testServerHandler{ + tcp: func(context.Context, net.Conn, M.Metadata) error { return errors.New("unexpected TCP") }, + udp: func(_ context.Context, conn N.PacketConn, _ M.Metadata) error { + calls.Add(1) + for index := 0; index < 3; index++ { + packet := buf.NewSize(MaxPayloadSize) + destination, err := conn.ReadPacket(packet) + if err != nil { + packet.Release() + return err + } + err = conn.WritePacket(packet, destination) + packet.Release() + if err != nil { + return err + } + } + return nil + }, + } + globalID := [8]byte{8, 7, 6, 5, 4, 3, 2, 1} + first := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, first, Frame{ + SessionID: 1, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "stable.example", Port: 53, GlobalID: globalID, Payload: []byte("one"), + }) + if response := readTestFrame(t, first); string(response.Payload) != "one" { + t.Fatalf("first response = %+v", response) + } + + second := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, second, Frame{ + SessionID: 2, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "stable.example", Port: 53, GlobalID: globalID, Payload: []byte("two"), + }) + if retired := readTestFrame(t, first); retired.Status != StatusEnd || retired.SessionID != 1 { + t.Fatalf("retired attachment frame = %+v", retired) + } + if response := readTestFrame(t, second); response.SessionID != 2 || string(response.Payload) != "two" { + t.Fatalf("second response = %+v", response) + } + + // A late frame from the retired carrier must not detach generation 2. + writeTestFrame(t, first, Frame{SessionID: 1, Status: StatusEnd}) + writeTestFrame(t, second, Frame{ + SessionID: 2, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, + Destination: "stable.example", Port: 53, Payload: []byte("three"), + }) + if response := readTestFrame(t, second); response.SessionID != 2 || string(response.Payload) != "three" { + t.Fatalf("post-stale response = %+v", response) + } + if got := calls.Load(); got != 1 { + t.Fatalf("XUDP dispatch calls = %d, want 1", got) + } +} + +func TestServerRuntimeRejectsXUDPTargetMismatchWithoutReplacingFlow(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{XUDPIdleTimeout: time.Hour}) + t.Cleanup(func() { _ = runtime.Close() }) + var calls atomic.Int32 + release := make(chan struct{}) + handler := testServerHandler{ + tcp: func(context.Context, net.Conn, M.Metadata) error { return errors.New("unexpected TCP") }, + udp: func(_ context.Context, conn N.PacketConn, _ M.Metadata) error { + calls.Add(1) + packet := buf.NewSize(MaxPayloadSize) + defer packet.Release() + _, err := conn.ReadPacket(packet) + if err != nil { + return err + } + <-release + return nil + }, + } + globalID := [8]byte{4, 4, 4, 4, 4, 4, 4, 4} + first := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, first, Frame{ + SessionID: 1, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "first.example", Port: 53, GlobalID: globalID, Payload: []byte("one"), + }) + writeTestFrame(t, first, Frame{SessionID: 1, Status: StatusEnd}) + + second := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, second, Frame{ + SessionID: 2, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "other.example", Port: 53, GlobalID: globalID, Payload: []byte("two"), + }) + response := readTestFrame(t, second) + if response.SessionID != 2 || response.Status != StatusEnd || response.Option&OptionError == 0 { + t.Fatalf("target mismatch response = %+v", response) + } + if got := calls.Load(); got != 1 { + t.Fatalf("XUDP dispatch calls = %d, want 1", got) + } + close(release) +} + +func TestServerRuntimeIsolatesXUDPFlowsByAuthenticatedUser(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{XUDPIdleTimeout: time.Hour}) + t.Cleanup(func() { _ = runtime.Close() }) + var calls atomic.Int32 + release := make(chan struct{}) + handler := testServerHandler{ + tcp: func(context.Context, net.Conn, M.Metadata) error { return errors.New("unexpected TCP") }, + udp: func(_ context.Context, conn N.PacketConn, _ M.Metadata) error { + calls.Add(1) + packet := buf.NewSize(MaxPayloadSize) + defer packet.Release() + destination, err := conn.ReadPacket(packet) + if err != nil { + return err + } + if err := conn.WritePacket(packet, destination); err != nil { + return err + } + <-release + return nil + }, + } + globalID := [8]byte{9, 9, 9, 9, 9, 9, 9, 9} + users := []string{"alice", "bob"} + for index, user := range users { + ctx := auth.ContextWithUser(context.Background(), user) + carrier := serveTestCarrierContext(t, ctx, runtime, handler) + writeTestFrame(t, carrier, Frame{ + SessionID: uint16(index + 1), Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "isolated.example", Port: 53, GlobalID: globalID, Payload: []byte(user), + }) + if response := readTestFrame(t, carrier); string(response.Payload) != user { + t.Fatalf("%s response = %+v", user, response) + } + } + if got := calls.Load(); got != 2 { + t.Fatalf("XUDP dispatch calls = %d, want 2", got) + } + close(release) +} + +func TestServerRuntimeIgnoresStaleXUDPExpiryAfterRebind(t *testing.T) { + timers := make(chan *manualServerTimer, 1) + runtime := NewServerRuntime(ServerOptions{ + XUDPIdleTimeout: time.Minute, + AfterFunc: func(_ time.Duration, callback func()) ServerTimer { + timer := &manualServerTimer{callback: callback} + timers <- timer + return timer + }, + }) + t.Cleanup(func() { _ = runtime.Close() }) + var calls atomic.Int32 + handler := testServerHandler{ + tcp: func(context.Context, net.Conn, M.Metadata) error { return errors.New("unexpected TCP") }, + udp: func(_ context.Context, conn N.PacketConn, _ M.Metadata) error { + calls.Add(1) + for index := 0; index < 3; index++ { + packet := buf.NewSize(MaxPayloadSize) + destination, err := conn.ReadPacket(packet) + if err != nil { + packet.Release() + return err + } + err = conn.WritePacket(packet, destination) + packet.Release() + if err != nil { + return err + } + } + return nil + }, + } + globalID := [8]byte{3, 1, 4, 1, 5, 9, 2, 6} + first := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, first, Frame{ + SessionID: 1, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "expiry.example", Port: 53, GlobalID: globalID, Payload: []byte("one"), + }) + if response := readTestFrame(t, first); string(response.Payload) != "one" { + t.Fatalf("first response = %+v", response) + } + writeTestFrame(t, first, Frame{SessionID: 1, Status: StatusEnd}) + staleTimer := <-timers + + second := serveTestCarrier(t, runtime, handler) + writeTestFrame(t, second, Frame{ + SessionID: 2, Status: StatusNew, Option: OptionData, Network: NetworkUDP, + Destination: "expiry.example", Port: 53, GlobalID: globalID, Payload: []byte("two"), + }) + if response := readTestFrame(t, second); string(response.Payload) != "two" { + t.Fatalf("second response = %+v", response) + } + + // Simulate a callback that had already raced with Stop(). + staleTimer.Fire() + writeTestFrame(t, second, Frame{ + SessionID: 2, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, + Destination: "expiry.example", Port: 53, Payload: []byte("three"), + }) + if response := readTestFrame(t, second); string(response.Payload) != "three" { + t.Fatalf("post-expiry response = %+v", response) + } + if got := calls.Load(); got != 1 { + t.Fatalf("XUDP dispatch calls = %d, want 1", got) + } +} + +func TestServerXUDPStaleWriteFailureDoesNotCloseCurrentGeneration(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{}) + ctx, cancel := context.WithCancelCause(context.Background()) + flow := newServerPacketFlow( + runtime, + xudpFlowKey{principal: "user:test", globalID: [8]byte{1}}, + ctx, + cancel, + testServerHandler{}, + testCarrierMetadata().Source, + M.Socksaddr{Fqdn: "write.example", Port: 53}, + true, + ) + oldConn := &blockingErrorConn{started: make(chan struct{}), release: make(chan struct{})} + oldCarrier := newServerCarrier(runtime, context.Background(), oldConn, testCarrierMetadata(), testServerHandler{}) + oldAttachment := &serverPacketAttachment{flow: flow, carrier: oldCarrier, id: 1, generation: 1} + flow.current = oldAttachment + flow.currentFast.Store(oldAttachment) + flow.generation = 1 + + packet := buf.NewSize(MaxPayloadSize) + if _, err := packet.Write([]byte("response")); err != nil { + t.Fatal(err) + } + defer packet.Release() + result := make(chan error, 1) + go func() { + result <- flow.WritePacket(packet, M.Socksaddr{Fqdn: "write.example", Port: 53}) + }() + <-oldConn.started + + newClient, newServer := net.Pipe() + defer newClient.Close() + defer newServer.Close() + newCarrier := newServerCarrier(runtime, context.Background(), newServer, testCarrierMetadata(), testServerHandler{}) + newAttachment := &serverPacketAttachment{flow: flow, carrier: newCarrier, id: 2, generation: 2} + flow.mu.Lock() + flow.current = newAttachment + flow.currentFast.Store(newAttachment) + flow.generation = 2 + flow.mu.Unlock() + close(oldConn.release) + + if err := <-result; err != nil { + t.Fatalf("stale write error = %v, want suppression", err) + } + flow.mu.Lock() + current := flow.current + flow.mu.Unlock() + if current != newAttachment { + t.Fatal("stale write changed current attachment") + } +} + +func TestServerRuntimeCloseDrainsCarriers(t *testing.T) { + runtime := NewServerRuntime(ServerOptions{XUDPIdleTimeout: time.Hour}) + client, server := net.Pipe() + done := make(chan error, 1) + handler := testServerHandler{ + tcp: func(context.Context, net.Conn, M.Metadata) error { return nil }, + udp: func(context.Context, N.PacketConn, M.Metadata) error { return nil }, + } + go func() { done <- runtime.Serve(context.Background(), server, testCarrierMetadata(), handler) }() + + if err := runtime.Close(); err != nil { + t.Fatal(err) + } + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("carrier did not drain") + } + _ = client.Close() + if err := runtime.Serve(context.Background(), client, testCarrierMetadata(), handler); !errors.Is(err, ErrServerClosed) { + t.Fatalf("Serve after Close = %v, want ErrServerClosed", err) + } +} diff --git a/transport/muxcool/session.go b/transport/muxcool/session.go new file mode 100644 index 0000000000..0aca58c9c5 --- /dev/null +++ b/transport/muxcool/session.go @@ -0,0 +1,310 @@ +package muxcool + +import ( + "context" + "errors" + "io" + "net" + "strconv" + "sync" + "time" +) + +type sessionOwner interface { + // writeFrame consumes frame and its payload before returning. + writeFrame(Frame) error + removeSession(uint16) +} + +type workerSession interface { + deliverDecodedFrame(decodedFrame) error + closeCarrier(error) +} + +// logicalConn is the caller-facing half of a logical Mux stream. net.Pipe gives +// each direction independent buffering and deadline state while the session +// goroutines translate the other half to and from mux.cool frames. +type logicalConn struct { + net.Conn + localAddr net.Addr + remoteAddr net.Addr + session *session +} + +func (c *logicalConn) LocalAddr() net.Addr { return c.localAddr } +func (c *logicalConn) RemoteAddr() net.Addr { return c.remoteAddr } + +func (c *logicalConn) Read(p []byte) (int, error) { + n, err := c.Conn.Read(p) + if err != nil { + if cause := c.session.terminalCause(); cause != nil { + return n, cause + } + } + return n, err +} + +func (c *logicalConn) Write(p []byte) (int, error) { + n, err := c.Conn.Write(p) + if err != nil { + if cause := c.session.terminalCause(); cause != nil { + return n, cause + } + } + return n, err +} + +type muxAddr string + +func (a muxAddr) Network() string { return "mux.cool" } +func (a muxAddr) String() string { return string(a) } + +type muxRemoteAddr struct { + host string + port uint16 +} + +func (muxRemoteAddr) Network() string { return "mux.cool" } +func (a muxRemoteAddr) String() string { + return net.JoinHostPort(a.host, strconv.Itoa(int(a.port))) +} + +type session struct { + owner sessionOwner + id uint16 + destination string + port uint16 + peer net.Conn + client net.Conn + downlink chan downlinkMessage + done chan struct{} + closeOnce sync.Once + causeMu sync.Mutex + cause error +} + +type downlinkMessage struct { + payload []byte + payloadPooled bool + terminal bool + cause error +} + +func (m *downlinkMessage) releasePayload() { + releasePooledPayload(m.payload, m.payloadPooled) + m.payload = nil + m.payloadPooled = false +} + +var sessionBufferPool = sync.Pool{ + New: func() any { return new([MaxPayloadSize]byte) }, +} + +func newSession( + ctx context.Context, + owner sessionOwner, + id uint16, + destination string, + port uint16, + firstPayloadTimeout time.Duration, +) (net.Conn, *session) { + logical, s := makeSession(owner, id, destination, port) + s.start(ctx, firstPayloadTimeout) + return logical, s +} + +func makeSession(owner sessionOwner, id uint16, destination string, port uint16) (net.Conn, *session) { + client, peer := net.Pipe() + s := &session{ + owner: owner, + id: id, + destination: destination, + port: port, + peer: peer, + client: client, + downlink: make(chan downlinkMessage, 16), + done: make(chan struct{}), + } + logical := &logicalConn{ + Conn: client, + localAddr: muxAddr("mux.cool"), + remoteAddr: muxRemoteAddr{host: destination, port: port}, + session: s, + } + return logical, s +} + +func (s *session) start(ctx context.Context, firstPayloadTimeout time.Duration) { + go s.runUplink(firstPayloadTimeout) + go s.runDownlink() + ctxDone := ctx.Done() + if ctxDone == nil { + return + } + go func() { + select { + case <-ctxDone: + s.finish(context.Cause(ctx), true) + case <-s.done: + } + }() +} + +func (s *session) runUplink(firstPayloadTimeout time.Duration) { + if firstPayloadTimeout > 0 { + _ = s.peer.SetReadDeadline(time.Now().Add(firstPayloadTimeout)) + } + + sentNew := false + pooledBuffer := sessionBufferPool.Get().(*[MaxPayloadSize]byte) + defer sessionBufferPool.Put(pooledBuffer) + buffer := pooledBuffer[:] + for { + n, err := s.peer.Read(buffer) + if err != nil && !sentNew && isNetTimeout(err) { + if writeErr := s.owner.writeFrame(Frame{ + SessionID: s.id, + Status: StatusNew, + Network: NetworkTCP, + Destination: s.destination, + Port: s.port, + }); writeErr != nil { + s.finish(writeErr, false) + return + } + sentNew = true + _ = s.peer.SetReadDeadline(time.Time{}) + continue + } + if n > 0 { + frame := Frame{ + SessionID: s.id, + Status: StatusKeep, + Option: OptionData, + Payload: buffer[:n], + } + if !sentNew { + frame.Status = StatusNew + frame.Network = NetworkTCP + frame.Destination = s.destination + frame.Port = s.port + sentNew = true + _ = s.peer.SetReadDeadline(time.Time{}) + } + if writeErr := s.owner.writeFrame(frame); writeErr != nil { + s.finish(writeErr, false) + return + } + } + if err != nil { + s.finish(nil, errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed)) + return + } + } +} + +func (s *session) runDownlink() { + for { + select { + case message := <-s.downlink: + err := writeFull(s.peer, message.payload) + message.releasePayload() + if err != nil { + s.finish(nil, true) + return + } + if message.terminal { + s.finish(message.cause, false) + return + } + case <-s.done: + return + } + } +} + +func (s *session) deliver(payload []byte) error { + if len(payload) == 0 { + return nil + } + // DecodeFrame allocates payload ownership for this session. Transfer that + // buffer directly to the downlink queue instead of copying every frame. + return s.enqueueDownlink(downlinkMessage{payload: payload}) +} + +func (s *session) deliverFinal(payload []byte, cause error) error { + return s.enqueueDownlink(downlinkMessage{ + payload: payload, + terminal: true, + cause: cause, + }) +} + +func (s *session) deliverFrame(frame Frame) error { + return s.deliverDecodedFrame(decodedFrame{Frame: frame}) +} + +func (s *session) deliverDecodedFrame(decoded decodedFrame) error { + frame := decoded.Frame + terminal := frame.Status == StatusEnd || frame.Option&OptionError != 0 + message := downlinkMessage{payload: frame.Payload, payloadPooled: decoded.payloadPooled} + if terminal { + var cause error + if frame.Option&OptionError != 0 { + cause = protocolError("remote session", errors.New("remote reported an error")) + } + message.terminal = true + message.cause = cause + return s.enqueueDownlink(message) + } + if frame.Option&OptionData != 0 { + return s.enqueueDownlink(message) + } + decoded.releasePayload() + return nil +} + +func (s *session) enqueueDownlink(message downlinkMessage) error { + select { + case s.downlink <- message: + return nil + case <-s.done: + message.releasePayload() + return net.ErrClosed + } +} + +func (s *session) closeCarrier(cause error) { + s.finish(cause, false) +} + +func (s *session) finish(cause error, sendEnd bool) { + s.closeOnce.Do(func() { + s.causeMu.Lock() + s.cause = cause + s.causeMu.Unlock() + if sendEnd { + _ = s.owner.writeFrame(Frame{SessionID: s.id, Status: StatusEnd}) + } + close(s.done) + _ = s.peer.Close() + // A clean remote End is represented by closing the peer side only. That + // lets the caller drain any final payload and then observe io.EOF instead + // of racing with a local close and receiving io.ErrClosedPipe. + if cause != nil || sendEnd { + _ = s.client.Close() + } + s.owner.removeSession(s.id) + }) +} + +func (s *session) terminalCause() error { + s.causeMu.Lock() + defer s.causeMu.Unlock() + return s.cause +} + +func isNetTimeout(err error) bool { + var netErr net.Error + return errors.As(err, &netErr) && netErr.Timeout() +} diff --git a/transport/muxcool/session_test.go b/transport/muxcool/session_test.go new file mode 100644 index 0000000000..5626a4235e --- /dev/null +++ b/transport/muxcool/session_test.go @@ -0,0 +1,185 @@ +package muxcool + +import ( + "context" + "errors" + "io" + "net" + "sync" + "testing" + "time" +) + +type fakeSessionOwner struct { + frames chan Frame + removed chan uint16 + err error + mu sync.Mutex +} + +func newFakeSessionOwner() *fakeSessionOwner { + return &fakeSessionOwner{ + frames: make(chan Frame, 32), + removed: make(chan uint16, 8), + } +} + +func (o *fakeSessionOwner) writeFrame(frame Frame) error { + o.mu.Lock() + err := o.err + o.mu.Unlock() + if err != nil { + return err + } + o.frames <- frame + return nil +} + +func (o *fakeSessionOwner) removeSession(id uint16) { + o.removed <- id +} + +func receiveFrame(t *testing.T, frames <-chan Frame) Frame { + t.Helper() + select { + case frame := <-frames: + return frame + case <-time.After(time.Second): + t.Fatal("timed out waiting for frame") + return Frame{} + } +} + +func TestLogicalSessionIsFullDuplexAndHasAddresses(t *testing.T) { + owner := newFakeSessionOwner() + conn, session := newSession(context.Background(), owner, 11, "echo.example", 443, time.Second) + t.Cleanup(func() { _ = conn.Close() }) + + writeDone := make(chan error, 1) + go func() { + _, err := conn.Write([]byte("upload")) + writeDone <- err + }() + frame := receiveFrame(t, owner.frames) + if frame.Status != StatusNew || string(frame.Payload) != "upload" { + t.Fatalf("uplink frame = status %d payload %q", frame.Status, frame.Payload) + } + if err := <-writeDone; err != nil { + t.Fatalf("write: %v", err) + } + + if err := session.deliver([]byte("download")); err != nil { + t.Fatalf("deliver: %v", err) + } + got := make([]byte, len("download")) + if _, err := io.ReadFull(conn, got); err != nil { + t.Fatalf("read: %v", err) + } + if string(got) != "download" { + t.Fatalf("download = %q", got) + } + if conn.LocalAddr() == nil || conn.RemoteAddr() == nil || conn.RemoteAddr().String() != "echo.example:443" { + t.Fatalf("addresses = local %v remote %v", conn.LocalAddr(), conn.RemoteAddr()) + } + if conn.LocalAddr().Network() != "mux.cool" || conn.RemoteAddr().Network() != "mux.cool" { + t.Fatalf("address networks = local %q remote %q", conn.LocalAddr().Network(), conn.RemoteAddr().Network()) + } +} + +func TestLogicalSessionDeadlines(t *testing.T) { + owner := newFakeSessionOwner() + conn, _ := newSession(context.Background(), owner, 12, "deadline.example", 80, time.Second) + t.Cleanup(func() { _ = conn.Close() }) + + if err := conn.SetReadDeadline(time.Now().Add(-time.Second)); err != nil { + t.Fatalf("SetReadDeadline: %v", err) + } + if _, err := conn.Read(make([]byte, 1)); !isTimeout(err) { + t.Fatalf("read error = %v, want timeout", err) + } + if err := conn.SetReadDeadline(time.Time{}); err != nil { + t.Fatalf("clear read deadline: %v", err) + } + if err := conn.SetWriteDeadline(time.Now().Add(-time.Second)); err != nil { + t.Fatalf("SetWriteDeadline: %v", err) + } + if _, err := conn.Write([]byte("x")); !isTimeout(err) { + t.Fatalf("write error = %v, want timeout", err) + } + if err := conn.SetDeadline(time.Time{}); err != nil { + t.Fatalf("clear deadline: %v", err) + } +} + +func isTimeout(err error) bool { + var netErr net.Error + return errors.As(err, &netErr) && netErr.Timeout() +} + +func TestLogicalSessionCancellationAndCloseAreIdempotent(t *testing.T) { + owner := newFakeSessionOwner() + ctx, cancel := context.WithCancel(context.Background()) + conn, _ := newSession(ctx, owner, 13, "cancel.example", 80, time.Second) + cancel() + + select { + case id := <-owner.removed: + if id != 13 { + t.Fatalf("removed session = %d", id) + } + case <-time.After(time.Second): + t.Fatal("session was not removed after cancellation") + } + if _, err := conn.Read(make([]byte, 1)); err == nil { + t.Fatal("read succeeded after cancellation") + } + if err := conn.Close(); err != nil { + t.Fatalf("first close: %v", err) + } + if err := conn.Close(); err != nil { + t.Fatalf("second close: %v", err) + } +} + +func TestSessionFrameStateMachine(t *testing.T) { + owner := newFakeSessionOwner() + conn, _ := newSession(context.Background(), owner, 21, "state.example", 8080, time.Second) + + if _, err := conn.Write([]byte("first")); err != nil { + t.Fatalf("first write: %v", err) + } + first := receiveFrame(t, owner.frames) + if first.Status != StatusNew || first.Destination != "state.example" || first.Port != 8080 { + t.Fatalf("first frame = %+v", first) + } + if _, err := conn.Write([]byte("second")); err != nil { + t.Fatalf("second write: %v", err) + } + second := receiveFrame(t, owner.frames) + if second.Status != StatusKeep || second.Destination != "" { + t.Fatalf("second frame = %+v", second) + } + if err := conn.Close(); err != nil { + t.Fatalf("close: %v", err) + } + end := receiveFrame(t, owner.frames) + if end.Status != StatusEnd || end.SessionID != 21 { + t.Fatalf("end frame = %+v", end) + } + select { + case extra := <-owner.frames: + t.Fatalf("unexpected extra frame: %+v", extra) + case <-time.After(20 * time.Millisecond): + } +} + +func TestSessionSendsPayloadFreeNewForServerFirstProtocol(t *testing.T) { + owner := newFakeSessionOwner() + conn, _ := newSession(context.Background(), owner, 22, "server-first.example", 25, 5*time.Millisecond) + t.Cleanup(func() { _ = conn.Close() }) + + frame := receiveFrame(t, owner.frames) + if frame.Status != StatusNew || frame.Option&OptionData != 0 || len(frame.Payload) != 0 { + t.Fatalf("payload-free new = %+v", frame) + } +} diff --git a/transport/muxcool/shutdown_test.go b/transport/muxcool/shutdown_test.go new file mode 100644 index 0000000000..b872ee9a67 --- /dev/null +++ b/transport/muxcool/shutdown_test.go @@ -0,0 +1,119 @@ +package muxcool + +import ( + "bytes" + "context" + "errors" + "io" + "net" + "sync" + "testing" + "time" +) + +func TestClosingOneSessionKeepsSiblingAlive(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + t.Cleanup(func() { worker.close(nil) }) + first, err := worker.openSession(context.Background(), "first.example", 80, time.Hour) + if err != nil { + t.Fatal(err) + } + second, err := worker.openSession(context.Background(), "second.example", 80, time.Hour) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = second.Close() }) + + if err := first.Close(); err != nil { + t.Fatal(err) + } + waitFor(t, func() bool { return worker.activeSessions() == 1 }) + carrier.inject(t, Frame{SessionID: 2, Status: StatusKeep, Option: OptionData, Payload: []byte("alive")}) + got := make([]byte, len("alive")) + if _, err := io.ReadFull(second, got); err != nil { + t.Fatalf("read sibling: %v", err) + } + if string(got) != "alive" { + t.Fatalf("sibling response = %q", got) + } +} + +func TestSessionContextAndRepeatedCloseSendOneEnd(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + t.Cleanup(func() { worker.close(nil) }) + ctx, cancel := context.WithCancel(context.Background()) + conn, err := worker.openSession(ctx, "cancel.example", 80, time.Hour) + if err != nil { + t.Fatal(err) + } + + cancel() + _ = conn.Close() + _ = conn.Close() + waitFor(t, func() bool { return worker.activeSessions() == 0 }) + + reader := bytes.NewReader(carrier.bytes()) + endCount := 0 + for reader.Len() > 0 { + frame, err := DecodeFrame(reader) + if err != nil { + t.Fatal(err) + } + if frame.SessionID == 1 && frame.Status == StatusEnd { + endCount++ + } + } + if endCount != 1 { + t.Fatalf("End count = %d, want 1", endCount) + } +} + +func TestSimultaneousSessionCarrierPoolAndContextShutdown(t *testing.T) { + dialer := &fakeCarrierDialer{} + options := testPoolOptions() + options.MaxConcurrency = 32 + pool := NewPool(dialer.dial, options) + + const sessionCount = 16 + connections := make([]net.Conn, 0, sessionCount) + cancels := make([]context.CancelFunc, 0, sessionCount) + for i := 0; i < sessionCount; i++ { + ctx, cancel := context.WithCancel(context.Background()) + conn, err := pool.DialContext(ctx, "shutdown.example", 443) + if err != nil { + t.Fatal(err) + } + connections = append(connections, conn) + cancels = append(cancels, cancel) + } + + var wg sync.WaitGroup + for index := range connections { + wg.Add(2) + go func(index int) { + defer wg.Done() + cancels[index]() + }(index) + go func(index int) { + defer wg.Done() + _ = connections[index].Close() + }(index) + } + wg.Add(3) + go func() { defer wg.Done(); _ = dialer.carriers[0].Close() }() + go func() { defer wg.Done(); _ = pool.Close() }() + go func() { defer wg.Done(); _ = pool.Close() }() + + done := make(chan struct{}) + go func() { wg.Wait(); close(done) }() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("simultaneous shutdown deadlocked") + } + if _, err := pool.DialContext(context.Background(), "closed.example", 80); !errors.Is(err, ErrPoolClosed) { + t.Fatalf("dial after close error = %v", err) + } +} diff --git a/transport/muxcool/worker.go b/transport/muxcool/worker.go new file mode 100644 index 0000000000..0212f9277b --- /dev/null +++ b/transport/muxcool/worker.go @@ -0,0 +1,218 @@ +package muxcool + +import ( + "context" + "errors" + "net" + "sync" + "sync/atomic" + "time" +) + +var errWorkerUnavailable = errors.New("mux.cool carrier is unavailable") + +type carrierWorker struct { + conn net.Conn + maxActive int + maxLifetime int + onClosed func(*carrierWorker) + onIdle func(*carrierWorker) + onAvailable func(*carrierWorker) + + writeMu sync.Mutex + writeBuffer []byte + mu sync.Mutex + sessions map[uint16]workerSession + nextID uint32 + lifetime int + draining bool + closed bool + closedFast atomic.Bool + closeErr error + closeOnce sync.Once +} + +func newCarrierWorker(conn net.Conn, maxActive, maxLifetime int, onClosed func(*carrierWorker)) *carrierWorker { + w := &carrierWorker{ + conn: conn, + maxActive: maxActive, + maxLifetime: maxLifetime, + onClosed: onClosed, + sessions: make(map[uint16]workerSession), + } + go w.readLoop() + return w +} + +func (w *carrierWorker) openSession( + ctx context.Context, + destination string, + port uint16, + firstPayloadTimeout time.Duration, +) (net.Conn, error) { + w.mu.Lock() + id, err := w.allocateIDLocked() + if err != nil { + w.mu.Unlock() + return nil, err + } + conn, logicalSession := makeSession(w, id, destination, port) + w.sessions[id] = logicalSession + w.mu.Unlock() + + logicalSession.start(ctx, firstPayloadTimeout) + return conn, nil +} + +func (w *carrierWorker) openPacketSession( + destination string, + port uint16, + globalID [8]byte, +) (net.PacketConn, error) { + w.mu.Lock() + id, err := w.allocateIDLocked() + if err != nil { + w.mu.Unlock() + return nil, err + } + logicalSession := makePacketSession(w, id, destination, port, globalID) + w.sessions[id] = logicalSession + w.mu.Unlock() + + return logicalSession, nil +} + +func (w *carrierWorker) allocateIDLocked() (uint16, error) { + if !w.availableLocked() { + err := w.closeErr + if err == nil { + err = errWorkerUnavailable + } + return 0, err + } + w.nextID++ + w.lifetime++ + if w.lifetime >= w.maxLifetime || w.nextID == uint32(^uint16(0)) { + w.draining = true + } + return uint16(w.nextID), nil +} + +func (w *carrierWorker) availableLocked() bool { + return !w.closed && !w.draining && len(w.sessions) < w.maxActive && w.lifetime < w.maxLifetime && w.nextID < uint32(^uint16(0)) +} + +func (w *carrierWorker) activeSessions() int { + w.mu.Lock() + defer w.mu.Unlock() + return len(w.sessions) +} + +func (w *carrierWorker) writeFrame(frame Frame) error { + w.writeMu.Lock() + raw, err := encodeFrame(w.writeBuffer, frame) + if err != nil { + w.writeMu.Unlock() + return err + } + w.writeBuffer = raw[:0] + if w.closedFast.Load() { + w.mu.Lock() + err = w.closeErr + if err == nil { + err = net.ErrClosed + } + w.mu.Unlock() + w.writeMu.Unlock() + return err + } + err = writeFull(w.conn, raw) + w.writeMu.Unlock() + if err != nil { + // Closing may re-enter writeFrame through a session's End path. Defer the + // fan-out until this call has returned to keep close paths acyclic. + go w.close(err) + } + return err +} + +func (w *carrierWorker) removeSession(id uint16) { + w.mu.Lock() + wasFull := len(w.sessions) >= w.maxActive + delete(w.sessions, id) + isIdle := len(w.sessions) == 0 + shouldClose := w.draining && isIdle + isAvailable := w.availableLocked() + w.mu.Unlock() + if shouldClose { + w.close(nil) + return + } + if wasFull && isAvailable && w.onAvailable != nil { + w.onAvailable(w) + } + if isIdle && w.onIdle != nil { + w.onIdle(w) + } +} + +func (w *carrierWorker) readLoop() { + metadataBuffer := make([]byte, MaxMetadataSize) + for { + frame, err := decodeFramePooled(w.conn, metadataBuffer) + if err != nil { + w.close(err) + return + } + if frame.Status == StatusKeepAlive { + frame.releasePayload() + continue + } + + w.mu.Lock() + logicalSession := w.sessions[frame.SessionID] + w.mu.Unlock() + if logicalSession == nil { + frame.releasePayload() + if frame.Status == StatusEnd { + continue + } + if err := w.writeFrame(Frame{SessionID: frame.SessionID, Status: StatusEnd}); err != nil { + return + } + continue + } + + if err := logicalSession.deliverDecodedFrame(frame); err != nil && !errors.Is(err, net.ErrClosed) { + w.close(err) + return + } + } +} + +func (w *carrierWorker) close(cause error) { + w.closeOnce.Do(func() { + if cause == nil { + cause = net.ErrClosed + } + + w.closedFast.Store(true) + w.mu.Lock() + w.closed = true + w.closeErr = cause + sessions := make([]workerSession, 0, len(w.sessions)) + for _, logicalSession := range w.sessions { + sessions = append(sessions, logicalSession) + } + w.sessions = nil + w.mu.Unlock() + + _ = w.conn.Close() + for _, logicalSession := range sessions { + logicalSession.closeCarrier(cause) + } + if w.onClosed != nil { + w.onClosed(w) + } + }) +} diff --git a/transport/muxcool/worker_test.go b/transport/muxcool/worker_test.go new file mode 100644 index 0000000000..41e523c74c --- /dev/null +++ b/transport/muxcool/worker_test.go @@ -0,0 +1,285 @@ +package muxcool + +import ( + "bytes" + "context" + "errors" + "io" + "net" + "sync" + "sync/atomic" + "testing" + "time" +) + +type testCarrier struct { + reader *io.PipeReader + serverWriter *io.PipeWriter + writesMu sync.Mutex + writes bytes.Buffer + activeWrites atomic.Int32 + interleaved atomic.Bool + writeErrMu sync.Mutex + writeErr error + closeOnce sync.Once + closed atomic.Bool +} + +func newTestCarrier() *testCarrier { + r, w := io.Pipe() + return &testCarrier{reader: r, serverWriter: w} +} + +func (c *testCarrier) Read(p []byte) (int, error) { return c.reader.Read(p) } + +func (c *testCarrier) Write(p []byte) (int, error) { + if c.activeWrites.Add(1) != 1 { + c.interleaved.Store(true) + } + defer c.activeWrites.Add(-1) + time.Sleep(time.Millisecond) + c.writeErrMu.Lock() + err := c.writeErr + c.writeErrMu.Unlock() + if err != nil { + return 0, err + } + c.writesMu.Lock() + defer c.writesMu.Unlock() + return c.writes.Write(p) +} + +func (c *testCarrier) Close() error { + c.closeOnce.Do(func() { + c.closed.Store(true) + _ = c.reader.Close() + _ = c.serverWriter.Close() + }) + return nil +} + +func (c *testCarrier) LocalAddr() net.Addr { return muxAddr("carrier-local") } +func (c *testCarrier) RemoteAddr() net.Addr { return muxAddr("carrier-remote") } +func (c *testCarrier) SetDeadline(time.Time) error { return nil } +func (c *testCarrier) SetReadDeadline(time.Time) error { return nil } +func (c *testCarrier) SetWriteDeadline(time.Time) error { return nil } + +func (c *testCarrier) setWriteError(err error) { + c.writeErrMu.Lock() + c.writeErr = err + c.writeErrMu.Unlock() +} + +func (c *testCarrier) bytes() []byte { + c.writesMu.Lock() + defer c.writesMu.Unlock() + return append([]byte(nil), c.writes.Bytes()...) +} + +func (c *testCarrier) isClosed() bool { + return c.closed.Load() +} + +func (c *testCarrier) inject(t *testing.T, frame Frame) { + t.Helper() + raw, err := EncodeFrame(frame) + if err != nil { + t.Fatal(err) + } + if _, err = c.serverWriter.Write(raw); err != nil { + t.Fatal(err) + } +} + +func TestCarrierWorkerSerializesConcurrentWriters(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 64, 128, nil) + t.Cleanup(func() { worker.close(nil) }) + + var wg sync.WaitGroup + for id := uint16(1); id <= 32; id++ { + wg.Add(1) + go func(id uint16) { + defer wg.Done() + if err := worker.writeFrame(Frame{SessionID: id, Status: StatusEnd}); err != nil { + t.Errorf("write frame %d: %v", id, err) + } + }(id) + } + wg.Wait() + if carrier.interleaved.Load() { + t.Fatal("carrier observed overlapping writes") + } + + reader := bytes.NewReader(carrier.bytes()) + for i := 0; i < 32; i++ { + if _, err := DecodeFrame(reader); err != nil { + t.Fatalf("decode serialized frames: %v", err) + } + } +} + +func TestCarrierWorkerDemultiplexesResponses(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + t.Cleanup(func() { worker.close(nil) }) + conn, err := worker.openSession(context.Background(), "echo.example", 443, time.Hour) + if err != nil { + t.Fatalf("open session: %v", err) + } + t.Cleanup(func() { _ = conn.Close() }) + + carrier.inject(t, Frame{SessionID: 1, Status: StatusKeep, Option: OptionData, Payload: []byte("response")}) + got := make([]byte, len("response")) + if _, err := io.ReadFull(conn, got); err != nil { + t.Fatalf("read response: %v", err) + } + if string(got) != "response" { + t.Fatalf("response = %q", got) + } +} + +func TestCarrierWorkerDemultiplexesStreamsAndPacketsOnOneCarrier(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + t.Cleanup(func() { worker.close(nil) }) + stream, err := worker.openSession(context.Background(), "stream.example", 443, time.Hour) + if err != nil { + t.Fatal(err) + } + packetConn, err := worker.openPacketSession("packet.example", 53, [8]byte{1}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = stream.Close(); _ = packetConn.Close() }) + + carrier.inject(t, Frame{SessionID: 1, Status: StatusKeep, Option: OptionData, Payload: []byte("stream")}) + carrier.inject(t, Frame{ + SessionID: 2, Status: StatusKeep, Option: OptionData, Network: NetworkUDP, + Destination: "9.9.9.9", Port: 53, Payload: []byte("packet"), + }) + + streamPayload := make([]byte, len("stream")) + if _, err := io.ReadFull(stream, streamPayload); err != nil { + t.Fatal(err) + } + packetPayload := make([]byte, len("packet")) + n, addr, err := packetConn.ReadFrom(packetPayload) + if err != nil { + t.Fatal(err) + } + if string(streamPayload) != "stream" || n != len(packetPayload) || string(packetPayload) != "packet" || addr.String() != "9.9.9.9:53" { + t.Fatalf("responses = stream %q, packet (%d, %q, %v)", streamPayload, n, packetPayload, addr) + } +} + +func TestCarrierWorkerEndsUnknownSession(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + t.Cleanup(func() { worker.close(nil) }) + + carrier.inject(t, Frame{SessionID: 77, Status: StatusKeep, Option: OptionData, Payload: []byte("orphan")}) + deadline := time.Now().Add(time.Second) + for len(carrier.bytes()) == 0 && time.Now().Before(deadline) { + time.Sleep(time.Millisecond) + } + frame, err := DecodeFrame(bytes.NewReader(carrier.bytes())) + if err != nil { + t.Fatalf("decode End: %v", err) + } + if frame.SessionID != 77 || frame.Status != StatusEnd { + t.Fatalf("unknown-session response = %+v", frame) + } +} + +func TestCarrierWorkerClosesSessionsOnMalformedFrameAndEOF(t *testing.T) { + for _, tc := range []struct { + name string + fail func(*testCarrier) + }{ + {name: "malformed", fail: func(c *testCarrier) { _, _ = c.serverWriter.Write([]byte{0, 3, 0, 1, 2}) }}, + {name: "EOF", fail: func(c *testCarrier) { _ = c.serverWriter.Close() }}, + } { + t.Run(tc.name, func(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + conn, err := worker.openSession(context.Background(), "failure.example", 80, time.Hour) + if err != nil { + t.Fatalf("open session: %v", err) + } + tc.fail(carrier) + _ = conn.SetReadDeadline(time.Now().Add(time.Second)) + if _, err := conn.Read(make([]byte, 1)); err == nil { + t.Fatal("logical session remained open") + } + }) + } +} + +func TestCarrierWorkerPropagatesWriteFailure(t *testing.T) { + carrierErr := errors.New("carrier write failed") + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + conn, err := worker.openSession(context.Background(), "failure.example", 80, time.Hour) + if err != nil { + t.Fatalf("open session: %v", err) + } + carrier.setWriteError(carrierErr) + _, _ = conn.Write([]byte("request")) + _ = conn.SetReadDeadline(time.Now().Add(time.Second)) + _, err = conn.Read(make([]byte, 1)) + if !errors.Is(err, carrierErr) { + t.Fatalf("read error = %v, want %v", err, carrierErr) + } +} + +func TestCarrierWorkerHandlesRemoteEndAndError(t *testing.T) { + tests := []struct { + name string + option Option + wantTyped bool + }{ + {name: "remote end"}, + {name: "remote error", option: OptionError, wantTyped: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + t.Cleanup(func() { worker.close(nil) }) + conn, err := worker.openSession(context.Background(), "remote.example", 80, time.Hour) + if err != nil { + t.Fatal(err) + } + carrier.inject(t, Frame{SessionID: 1, Status: StatusEnd, Option: tt.option}) + _ = conn.SetReadDeadline(time.Now().Add(time.Second)) + _, err = conn.Read(make([]byte, 1)) + if err == nil { + t.Fatal("read succeeded after remote close") + } + var protocolErr *ProtocolError + if errors.As(err, &protocolErr) != tt.wantTyped { + t.Fatalf("read error = %v, typed protocol error = %t", err, errors.As(err, &protocolErr)) + } + waitFor(t, func() bool { return worker.activeSessions() == 0 }) + }) + } +} + +func TestCarrierWorkerDeliversFinalPayloadBeforeRemoteEnd(t *testing.T) { + carrier := newTestCarrier() + worker := newCarrierWorker(carrier, 8, 128, nil) + t.Cleanup(func() { worker.close(nil) }) + conn, err := worker.openSession(context.Background(), "final.example", 80, time.Hour) + if err != nil { + t.Fatal(err) + } + carrier.inject(t, Frame{SessionID: 1, Status: StatusEnd, Option: OptionData, Payload: []byte("final")}) + got, err := io.ReadAll(conn) + if err != nil { + t.Fatalf("read final payload: %v", err) + } + if string(got) != "final" { + t.Fatalf("final payload = %q", got) + } +} diff --git a/transport/muxcool/xray_process_e2e_test.go b/transport/muxcool/xray_process_e2e_test.go new file mode 100644 index 0000000000..43bbc28658 --- /dev/null +++ b/transport/muxcool/xray_process_e2e_test.go @@ -0,0 +1,498 @@ +//go:build integration + +package muxcool_test + +import ( + "bytes" + "context" + "fmt" + "net" + "os" + "os/exec" + "path/filepath" + "runtime" + "strings" + "testing" + "time" + + "github.com/metacubex/mihomo/transport/socks5" + "github.com/miekg/dns" +) + +func TestMuxCoolProcessXrayCompatibility(t *testing.T) { + mihomoBinary := buildMihomoProcessE2E(t) + xrayBinary := buildXrayProcessE2E(t) + + t.Run("mihomo-client/xray-server/tcp", func(t *testing.T) { + tcpAddress, _ := startProcessE2EEcho(t) + serverPort := reserveProcessE2EPort(t) + clientPort := reserveProcessE2EPort(t) + tempDir := t.TempDir() + + serverConfig := filepath.Join(tempDir, "xray-server.json") + clientConfig := filepath.Join(tempDir, "mihomo-client.yaml") + writeProcessE2EFile(t, serverConfig, xrayServerProcessE2EConfig(serverPort)) + writeProcessE2EFile(t, clientConfig, mihomoClientProcessE2EConfig(clientPort, serverPort, 0)) + + server := startXrayProcessE2E(t, xrayBinary, serverConfig, nil) + waitProcessE2ETCP(t, server, fmt.Sprintf("127.0.0.1:%d", serverPort)) + client := startMihomoProcessE2E(t, mihomoBinary, clientConfig, filepath.Join(tempDir, "mihomo-home")) + clientAddress := fmt.Sprintf("127.0.0.1:%d", clientPort) + waitProcessE2ETCP(t, client, clientAddress) + + processE2ETCPRoundTrip(t, clientAddress, tcpAddress) + }) + + t.Run("xray-client/mihomo-server/tcp", func(t *testing.T) { + tcpAddress, _ := startProcessE2EEcho(t) + serverPort := reserveProcessE2EPort(t) + clientPort := reserveProcessE2EPort(t) + tempDir := t.TempDir() + + serverConfig := filepath.Join(tempDir, "mihomo-server.yaml") + clientConfig := filepath.Join(tempDir, "xray-client.json") + writeProcessE2EFile(t, serverConfig, mihomoServerProcessE2EConfig(serverPort)) + writeProcessE2EFile(t, clientConfig, xrayClientProcessE2EConfig(clientPort, serverPort, 0)) + + server := startMihomoProcessE2E(t, mihomoBinary, serverConfig, filepath.Join(tempDir, "mihomo-home")) + waitProcessE2ETCP(t, server, fmt.Sprintf("127.0.0.1:%d", serverPort)) + client := startXrayProcessE2E(t, xrayBinary, clientConfig, []string{"XRAY_CONE_DISABLED=false"}) + clientAddress := fmt.Sprintf("127.0.0.1:%d", clientPort) + waitProcessE2ETCP(t, client, clientAddress) + + processE2ETCPRoundTrip(t, clientAddress, tcpAddress) + }) + + t.Run("mihomo-client/xray-server/udp", func(t *testing.T) { + dnsTarget := startProcessE2EDNSServer(t) + serverPort := reserveProcessE2EPort(t) + clientPort := reserveProcessE2EPort(t) + dnsPort := reserveProcessE2EPort(t) + tempDir := t.TempDir() + + serverConfig := filepath.Join(tempDir, "xray-server.json") + clientConfig := filepath.Join(tempDir, "mihomo-client.yaml") + writeProcessE2EFile(t, serverConfig, xrayServerProcessE2EConfig(serverPort)) + writeProcessE2EFile(t, clientConfig, mihomoDNSClientProcessE2EConfig(clientPort, dnsPort, serverPort, dnsTarget)) + + server := startXrayProcessE2E(t, xrayBinary, serverConfig, nil) + waitProcessE2ETCP(t, server, fmt.Sprintf("127.0.0.1:%d", serverPort)) + client := startMihomoProcessE2E(t, mihomoBinary, clientConfig, filepath.Join(tempDir, "mihomo-home")) + waitProcessE2ETCP(t, client, fmt.Sprintf("127.0.0.1:%d", clientPort)) + + processE2EDNSRoundTrip(t, fmt.Sprintf("127.0.0.1:%d", dnsPort)) + }) + + t.Run("mihomo-client/xray-server/xudp", func(t *testing.T) { + udpAddress := startProcessE2EUDPEcho(t) + serverPort := reserveProcessE2EPort(t) + clientPort := reserveProcessE2EPort(t) + tempDir := t.TempDir() + + serverConfig := filepath.Join(tempDir, "xray-server.json") + clientConfig := filepath.Join(tempDir, "mihomo-client.yaml") + writeProcessE2EFile(t, serverConfig, xrayServerProcessE2EConfig(serverPort)) + writeProcessE2EFile(t, clientConfig, mihomoClientProcessE2EConfig(clientPort, serverPort, 4)) + + server := startXrayProcessE2E(t, xrayBinary, serverConfig, nil) + waitProcessE2ETCP(t, server, fmt.Sprintf("127.0.0.1:%d", serverPort)) + client := startMihomoProcessE2E(t, mihomoBinary, clientConfig, filepath.Join(tempDir, "mihomo-home")) + clientAddress := fmt.Sprintf("127.0.0.1:%d", clientPort) + waitProcessE2ETCP(t, client, clientAddress) + + processE2ESOCKSUDPRoundTrip(t, clientAddress, udpAddress) + }) + + t.Run("xray-client/mihomo-server/udp", func(t *testing.T) { + udpAddress := startProcessE2EUDPEcho(t) + serverPort := reserveProcessE2EPort(t) + clientPort := reserveProcessE2EPort(t) + tempDir := t.TempDir() + + serverConfig := filepath.Join(tempDir, "mihomo-server.yaml") + clientConfig := filepath.Join(tempDir, "xray-client.json") + writeProcessE2EFile(t, serverConfig, mihomoServerProcessE2EConfig(serverPort)) + writeProcessE2EFile(t, clientConfig, xrayClientProcessE2EConfig(clientPort, serverPort, 0)) + + server := startMihomoProcessE2E(t, mihomoBinary, serverConfig, filepath.Join(tempDir, "mihomo-home")) + waitProcessE2ETCP(t, server, fmt.Sprintf("127.0.0.1:%d", serverPort)) + client := startXrayProcessE2E(t, xrayBinary, clientConfig, []string{"XRAY_CONE_DISABLED=true"}) + clientAddress := fmt.Sprintf("127.0.0.1:%d", clientPort) + waitProcessE2ETCP(t, client, clientAddress) + + processE2ESOCKSUDPRoundTrip(t, clientAddress, udpAddress) + }) + + t.Run("xray-client/mihomo-server/xudp", func(t *testing.T) { + udpAddress := startProcessE2EUDPEcho(t) + serverPort := reserveProcessE2EPort(t) + clientPort := reserveProcessE2EPort(t) + tempDir := t.TempDir() + + serverConfig := filepath.Join(tempDir, "mihomo-server.yaml") + clientConfig := filepath.Join(tempDir, "xray-client.json") + writeProcessE2EFile(t, serverConfig, mihomoServerProcessE2EConfig(serverPort)) + writeProcessE2EFile(t, clientConfig, xrayClientProcessE2EConfig(clientPort, serverPort, 4)) + + server := startMihomoProcessE2E(t, mihomoBinary, serverConfig, filepath.Join(tempDir, "mihomo-home")) + waitProcessE2ETCP(t, server, fmt.Sprintf("127.0.0.1:%d", serverPort)) + client := startXrayProcessE2E(t, xrayBinary, clientConfig, []string{"XRAY_CONE_DISABLED=false"}) + clientAddress := fmt.Sprintf("127.0.0.1:%d", clientPort) + waitProcessE2ETCP(t, client, clientAddress) + + processE2ESOCKSUDPRoundTrip(t, clientAddress, udpAddress) + }) +} + +func processE2ETCPRoundTrip(t *testing.T, clientAddress, targetAddress string) { + t.Helper() + deadline := time.Now().Add(10 * time.Second) + var lastError error + for time.Now().Before(deadline) { + result := <-startProcessE2EWave(clientAddress, targetAddress, 0, 1) + if result.err == nil { + _ = result.connection.Close() + return + } + lastError = result.err + time.Sleep(25 * time.Millisecond) + } + t.Fatalf("TCP round trip did not become ready: %v", lastError) +} + +func xrayServerProcessE2EConfig(port int) string { + return fmt.Sprintf(`{ + "log": {"loglevel": "debug"}, + "inbounds": [{ + "tag": "vless-in", + "listen": "127.0.0.1", + "port": %d, + "protocol": "vless", + "settings": { + "clients": [{"id": %q}], + "decryption": "none" + } + }], + "outbounds": [{ + "tag": "direct", + "protocol": "freedom", + "settings": {"finalRules": [{"action": "allow", "network": "tcp,udp"}]} + }] +}`, port, processE2EUUID) +} + +func xrayClientProcessE2EConfig(socksPort, serverPort, xudpConcurrency int) string { + return fmt.Sprintf(`{ + "log": {"loglevel": "debug"}, + "inbounds": [{ + "tag": "socks-in", + "listen": "127.0.0.1", + "port": %d, + "protocol": "socks", + "settings": {"auth": "noauth", "udp": true} + }], + "outbounds": [{ + "tag": "vless-mux-cool", + "protocol": "vless", + "settings": {"vnext": [{ + "address": "127.0.0.1", + "port": %d, + "users": [{"id": %q, "encryption": "none"}] + }]}, + "mux": { + "enabled": true, + "concurrency": 8, + "xudpConcurrency": %d, + "xudpProxyUDP443": "allow" + } + }] +}`, socksPort, serverPort, processE2EUUID, xudpConcurrency) +} + +func mihomoServerProcessE2EConfig(port int) string { + return fmt.Sprintf(` +mode: rule +log-level: debug +listeners: + - name: vless-in + type: vless + listen: 127.0.0.1 + port: %d + allow-insecure: true + users: + - username: process-e2e + uuid: %s +rules: + - MATCH,DIRECT +`, port, processE2EUUID) +} + +func mihomoClientProcessE2EConfig(mixedPort, serverPort, xudpConcurrency int) string { + return fmt.Sprintf(` +mixed-port: %d +allow-lan: false +mode: rule +log-level: debug +proxies: + - name: vless-mux-cool + type: vless + server: 127.0.0.1 + port: %d + uuid: %s + network: tcp + udp: true + mux.cool: + enabled: true + max-concurrency: 8 + max-connections: 128 + max-carriers: 4 + xudp-concurrency: %d + xudp-proxy-udp443: allow +rules: + - MATCH,vless-mux-cool +`, mixedPort, serverPort, processE2EUUID, xudpConcurrency) +} + +func mihomoDNSClientProcessE2EConfig(mixedPort, dnsPort, serverPort int, dnsTarget string) string { + return fmt.Sprintf(` +mixed-port: %d +allow-lan: false +mode: rule +log-level: debug +dns: + enable: true + listen: 127.0.0.1:%d + default-nameserver: + - 127.0.0.1 + nameserver: + - "udp://%s#vless-mux-cool" +proxies: + - name: vless-mux-cool + type: vless + server: 127.0.0.1 + port: %d + uuid: %s + network: tcp + udp: true + mux.cool: + enabled: true + max-concurrency: 8 + max-connections: 128 + max-carriers: 4 + xudp-concurrency: 0 + xudp-proxy-udp443: allow +rules: + - MATCH,vless-mux-cool +`, mixedPort, dnsPort, dnsTarget, serverPort, processE2EUUID) +} + +func startProcessE2EUDPEcho(t *testing.T) string { + t.Helper() + connection, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = connection.Close() }) + go func() { + buffer := make([]byte, 64*1024) + for { + length, source, err := connection.ReadFromUDP(buffer) + if err != nil { + return + } + _, _ = connection.WriteToUDP(buffer[:length], source) + } + }() + return connection.LocalAddr().String() +} + +func processE2ESOCKSUDPRoundTrip(t *testing.T, socksAddress, targetAddress string) { + t.Helper() + control, err := net.DialTimeout("tcp", socksAddress, 5*time.Second) + if err != nil { + t.Fatalf("dial SOCKS control connection: %v", err) + } + defer control.Close() + if err := control.SetDeadline(time.Now().Add(10 * time.Second)); err != nil { + t.Fatal(err) + } + relay, err := socks5.ClientHandshake(control, socks5.ParseAddr("0.0.0.0:0"), socks5.CmdUDPAssociate, nil) + if err != nil { + t.Fatalf("SOCKS5 UDP associate: %v", err) + } + + packetConnection, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatal(err) + } + defer packetConnection.Close() + relayAddress := relay.UDPAddr() + if relayAddress == nil { + t.Fatalf("SOCKS5 UDP relay address %q is not an IP address", relay) + } + if relayAddress.IP.IsUnspecified() { + relayAddress.IP = net.ParseIP("127.0.0.1") + } + + payload := []byte("mux.cool cross-runtime UDP payload") + target := socks5.ParseAddr(targetAddress) + packet, err := socks5.EncodeUDPPacket(target, payload) + if err != nil { + t.Fatal(err) + } + responseBuffer := make([]byte, 64*1024) + deadline := time.Now().Add(10 * time.Second) + var length int + for { + if _, err := packetConnection.WriteToUDP(packet, relayAddress); err != nil { + t.Fatalf("write SOCKS5 UDP packet: %v", err) + } + readDeadline := time.Now().Add(500 * time.Millisecond) + if readDeadline.After(deadline) { + readDeadline = deadline + } + if err := packetConnection.SetReadDeadline(readDeadline); err != nil { + t.Fatal(err) + } + length, _, err = packetConnection.ReadFromUDP(responseBuffer) + if err == nil { + break + } + if timeout, ok := err.(net.Error); !ok || !timeout.Timeout() || !time.Now().Before(deadline) { + t.Fatalf("read SOCKS5 UDP packet: %v", err) + } + } + responseTarget, responsePayload, err := socks5.DecodeUDPPacket(responseBuffer[:length]) + if err != nil { + t.Fatalf("decode SOCKS5 UDP packet: %v", err) + } + if responseTarget.String() != target.String() { + t.Fatalf("SOCKS5 UDP response target = %s, want %s", responseTarget, target) + } + if !bytes.Equal(responsePayload, payload) { + t.Fatalf("SOCKS5 UDP response payload = %q, want %q", responsePayload, payload) + } +} + +func startProcessE2EDNSServer(t *testing.T) string { + t.Helper() + connection, err := net.ListenUDP("udp", &net.UDPAddr{IP: net.ParseIP("127.0.0.1")}) + if err != nil { + t.Fatal(err) + } + server := &dns.Server{ + PacketConn: connection, + Handler: dns.HandlerFunc(func(writer dns.ResponseWriter, request *dns.Msg) { + response := new(dns.Msg) + response.SetReply(request) + response.Authoritative = true + if len(request.Question) > 0 { + response.Answer = append(response.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: request.Question[0].Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 60}, + A: net.ParseIP("192.0.2.1"), + }) + } + _ = writer.WriteMsg(response) + }), + } + t.Cleanup(func() { _ = server.Shutdown() }) + go func() { _ = server.ActivateAndServe() }() + return connection.LocalAddr().String() +} + +func processE2EDNSRoundTrip(t *testing.T, serverAddress string) { + t.Helper() + request := new(dns.Msg) + request.SetQuestion("mux-cool-process.test.", dns.TypeA) + response, _, err := (&dns.Client{Net: "udp", Timeout: 10 * time.Second}).Exchange(request, serverAddress) + if err != nil { + t.Fatalf("DNS round trip: %v", err) + } + if len(response.Answer) != 1 { + t.Fatalf("DNS answer count = %d, want 1", len(response.Answer)) + } + answer, ok := response.Answer[0].(*dns.A) + if !ok || !answer.A.Equal(net.ParseIP("192.0.2.1")) { + t.Fatalf("DNS answer = %v, want 192.0.2.1", response.Answer[0]) + } +} + +func startXrayProcessE2E(t *testing.T, binary, config string, environment []string) *mihomoProcessE2E { + t.Helper() + ctx, cancel := context.WithCancel(context.Background()) + process := &mihomoProcessE2E{cancel: cancel, done: make(chan struct{})} + command := exec.CommandContext(ctx, binary, "run", "-config", config) + command.Env = processE2EEnvironment(environment) + command.Stdout = &process.output + command.Stderr = &process.output + if err := command.Start(); err != nil { + cancel() + t.Fatal(err) + } + go func() { + err := command.Wait() + process.mu.Lock() + process.err = err + process.mu.Unlock() + close(process.done) + }() + t.Cleanup(func() { + cancel() + select { + case <-process.done: + case <-time.After(5 * time.Second): + t.Errorf("Xray process did not stop\n%s", process.output.String()) + } + if t.Failed() { + t.Log(process.output.String()) + } + }) + return process +} + +func processE2EEnvironment(overrides []string) []string { + environment := os.Environ() + for _, override := range overrides { + key, _, _ := strings.Cut(override, "=") + prefix := key + "=" + filtered := environment[:0] + for _, entry := range environment { + if !strings.HasPrefix(entry, prefix) { + filtered = append(filtered, entry) + } + } + environment = append(filtered, override) + } + return environment +} + +func buildXrayProcessE2E(t *testing.T) string { + t.Helper() + if binary := os.Getenv("XRAY_E2E_BINARY"); binary != "" { + return binary + } + + root := os.Getenv("XRAY_CORE_ROOT") + if root == "" { + mihomoRoot, err := filepath.Abs(filepath.Join("..", "..")) + if err != nil { + t.Fatal(err) + } + root = filepath.Join(filepath.Dir(mihomoRoot), "Xray-core") + } + if _, err := os.Stat(filepath.Join(root, "go.mod")); err != nil { + t.Skipf("Xray-core source is unavailable at %s; set XRAY_CORE_ROOT or XRAY_E2E_BINARY", root) + } + + binary := filepath.Join(t.TempDir(), "xray") + if runtime.GOOS == "windows" { + binary += ".exe" + } + command := exec.Command("go", "build", "-trimpath", "-o", binary, "./main") + command.Dir = root + output, err := command.CombinedOutput() + if err != nil { + t.Fatalf("build Xray: %v\n%s", err, output) + } + return binary +}