diff --git a/common/mux/client_test.go b/common/mux/client_test.go index 9626e2a27..79f0217e5 100644 --- a/common/mux/client_test.go +++ b/common/mux/client_test.go @@ -7,9 +7,11 @@ import ( "github.com/golang/mock/gomock" "github.com/xtls/xray-core/common" + "github.com/xtls/xray-core/common/buf" "github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/mux" "github.com/xtls/xray-core/common/net" + "github.com/xtls/xray-core/common/protocol" "github.com/xtls/xray-core/common/session" "github.com/xtls/xray-core/testing/mocks" "github.com/xtls/xray-core/transport" @@ -114,3 +116,48 @@ func TestClientWorkerClose(t *testing.T) { common.Must(w2.Close()) } + +func TestClientWorkerUDPSource(t *testing.T) { + downR, downW := pipe.New(pipe.WithoutSizeLimit()) + upR, upW := pipe.New(pipe.WithoutSizeLimit()) + worker, err := mux.NewClientWorker(transport.Link{Reader: downR, Writer: upW}, mux.ClientStrategy{}) + common.Must(err) + + inR, inW := pipe.New(pipe.WithoutSizeLimit()) + outR, outW := pipe.New(pipe.WithoutSizeLimit()) + ctx := session.ContextWithOutbounds(context.Background(), []*session.Outbound{{ + Target: net.UDPDestination(net.ParseAddress("8.8.8.8"), 53), + }}) + if !worker.Dispatch(ctx, &transport.Link{Reader: inR, Writer: outW}) { + t.Fatal("failed to dispatch") + } + b := buf.New() + b.WriteString("query") + common.Must(inW.WriteMultiBuffer(buf.MultiBuffer{b})) + mb, err := upR.ReadMultiBuffer() // New frame, the session is UDP from now on + common.Must(err) + buf.ReleaseMulti(mb) + + srcs := []net.Destination{ + net.UDPDestination(net.ParseAddress("1.1.1.1"), 1111), + net.UDPDestination(net.DomainAddress("example.com"), 2222), + net.UDPDestination(net.ParseAddress("3.3.3.3"), 3333), + } + w := mux.NewResponseWriter(1, downW, protocol.TransferTypePacket) + var got buf.MultiBuffer + for i := range srcs { + b := buf.New() + b.WriteString("reply") + b.UDP = &srcs[i] + common.Must(w.WriteMultiBuffer(buf.MultiBuffer{b})) + // keep earlier replies around while the next frame is parsed + mb, err := outR.ReadMultiBuffer() + common.Must(err) + got = append(got, mb...) + } + for i, b := range got { + if b.UDP == nil || *b.UDP != srcs[i] { + t.Errorf("reply %d: source = %v, want %v", i, b.UDP, srcs[i]) + } + } +} diff --git a/common/mux/reader.go b/common/mux/reader.go index b9714cdf9..697af1dee 100644 --- a/common/mux/reader.go +++ b/common/mux/reader.go @@ -14,15 +14,16 @@ import ( type PacketReader struct { reader io.Reader eof bool - dest *net.Destination + dest net.Destination } // NewPacketReader creates a new PacketReader. +// dest is copied because the caller reuses it for the next frame. func NewPacketReader(reader io.Reader, dest *net.Destination) *PacketReader { return &PacketReader{ reader: reader, eof: false, - dest: dest, + dest: *dest, } } @@ -47,8 +48,8 @@ func (r *PacketReader) ReadMultiBuffer() (buf.MultiBuffer, error) { return nil, err } r.eof = true - if r.dest != nil && r.dest.Network == net.Network_UDP { - b.UDP = r.dest + if r.dest.Network == net.Network_UDP { + b.UDP = &r.dest // only one packet is read, so b owns r.dest } return buf.MultiBuffer{b}, nil }