Compare commits

...
2 Commits
3 changed files with 61 additions and 5 deletions
+47
View File
@@ -7,9 +7,11 @@ import (
"github.com/golang/mock/gomock" "github.com/golang/mock/gomock"
"github.com/xtls/xray-core/common" "github.com/xtls/xray-core/common"
"github.com/xtls/xray-core/common/buf"
"github.com/xtls/xray-core/common/errors" "github.com/xtls/xray-core/common/errors"
"github.com/xtls/xray-core/common/mux" "github.com/xtls/xray-core/common/mux"
"github.com/xtls/xray-core/common/net" "github.com/xtls/xray-core/common/net"
"github.com/xtls/xray-core/common/protocol"
"github.com/xtls/xray-core/common/session" "github.com/xtls/xray-core/common/session"
"github.com/xtls/xray-core/testing/mocks" "github.com/xtls/xray-core/testing/mocks"
"github.com/xtls/xray-core/transport" "github.com/xtls/xray-core/transport"
@@ -114,3 +116,48 @@ func TestClientWorkerClose(t *testing.T) {
common.Must(w2.Close()) 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])
}
}
}
+5 -4
View File
@@ -14,15 +14,16 @@ import (
type PacketReader struct { type PacketReader struct {
reader io.Reader reader io.Reader
eof bool eof bool
dest *net.Destination dest net.Destination
} }
// NewPacketReader creates a new PacketReader. // 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 { func NewPacketReader(reader io.Reader, dest *net.Destination) *PacketReader {
return &PacketReader{ return &PacketReader{
reader: reader, reader: reader,
eof: false, eof: false,
dest: dest, dest: *dest,
} }
} }
@@ -47,8 +48,8 @@ func (r *PacketReader) ReadMultiBuffer() (buf.MultiBuffer, error) {
return nil, err return nil, err
} }
r.eof = true r.eof = true
if r.dest != nil && r.dest.Network == net.Network_UDP { if r.dest.Network == net.Network_UDP {
b.UDP = r.dest b.UDP = &r.dest // only one packet is read, so b owns r.dest
} }
return buf.MultiBuffer{b}, nil return buf.MultiBuffer{b}, nil
} }
+9 -1
View File
@@ -170,9 +170,17 @@ func (s *Server) processTCP(ctx context.Context, conn stat.Connection, dispatche
return errors.New("UDP associate with listen port failed") return errors.New("UDP associate with listen port failed")
} }
tempUDPConn.SetTimeout(plcy.Timeouts.ConnectionIdle) tempUDPConn.SetTimeout(plcy.Timeouts.ConnectionIdle)
var udpConn stat.Connection = tempUDPConn
if counters, ok := conn.(*stat.CounterConnection); ok {
udpConn = &stat.CounterConnection{
Connection: tempUDPConn,
ReadCounter: counters.ReadCounter,
WriteCounter: counters.WriteCounter,
}
}
errCh := make(chan error, 1) errCh := make(chan error, 1)
go func() { go func() {
errCh <- s.handleUDPPayload(ctx, tempUDPConn, dispatcher) errCh <- s.handleUDPPayload(ctx, udpConn, dispatcher)
}() }()
// Associated TCP keeps the UDP alive // Associated TCP keeps the UDP alive
// Close UDP if TCP connection is closed // Close UDP if TCP connection is closed