Skip to content

Commit 9b61266

Browse files
authored
sqlreplay, backend: fix bugs to make traffic replay runnable (#671)
1 parent 7a4a4c2 commit 9b61266

17 files changed

Lines changed: 342 additions & 91 deletions

File tree

pkg/proxy/backend/backend_conn_mgr.go

Lines changed: 3 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -321,7 +321,9 @@ func (mgr *BackendConnManager) ExecuteCmd(ctx context.Context, request []byte) (
321321
}
322322
mgr.processLock.Lock()
323323
defer func() {
324-
mgr.setQuitSourceByErr(err)
324+
if err != nil && !pnet.IsMySQLError(err) {
325+
mgr.setQuitSourceByErr(err)
326+
}
325327
mgr.handshakeHandler.OnTraffic(mgr)
326328
now := time.Now()
327329
if err != nil && errors.Is(err, ErrBackendConn) {
@@ -403,12 +405,7 @@ func (mgr *BackendConnManager) ExecuteCmd(ctx context.Context, request []byte) (
403405
_, err = mgr.cmdProcessor.executeCmd(request, mgr.clientIO, backendIO, false)
404406
addCmdMetrics(cmd, backendIO.RemoteAddr().String(), startTime)
405407
mgr.updateTraffic(backendIO)
406-
if err != nil && !pnet.IsMySQLError(err) {
407-
return
408-
}
409408
}
410-
// Ignore MySQL errors, only return unexpected errors.
411-
err = nil
412409
return
413410
}
414411

pkg/proxy/backend/backend_conn_mgr_test.go

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -199,6 +199,9 @@ func (ts *backendMgrTester) forwardCmd4Proxy(clientIO, backendIO pnet.PacketIO)
199199
prevCounter, err := readCmdCounter(pnet.Command(request[0]), ts.tc.backendListener.Addr().String())
200200
require.NoError(ts.t, err)
201201
rsErr := ts.mp.ExecuteCmd(context.Background(), request)
202+
if pnet.IsMySQLError(rsErr) {
203+
rsErr = nil
204+
}
202205
curCounter, err := readCmdCounter(pnet.Command(request[0]), ts.tc.backendListener.Addr().String())
203206
require.NoError(ts.t, err)
204207
require.Equal(ts.t, prevCounter+1, curCounter)
@@ -575,6 +578,41 @@ func TestSpecialCmds(t *testing.T) {
575578
ts.runTests(runners)
576579
}
577580

581+
// Test that ExecuteCmd may return a mysql error, which is required by traffic replay.
582+
func TestReturnMySQLError(t *testing.T) {
583+
ts := newBackendMgrTester(t)
584+
runners := []runner{
585+
// 1st handshake
586+
{
587+
client: ts.mc.authenticate,
588+
proxy: ts.firstHandshake4Proxy,
589+
backend: ts.handshake4Backend,
590+
},
591+
// mysql error
592+
{
593+
client: func(packetIO pnet.PacketIO) error {
594+
ts.mc.cmd = pnet.ComQuery
595+
ts.mc.sql = "select $$"
596+
return ts.mc.request(packetIO)
597+
},
598+
proxy: func(clientIO, backendIO pnet.PacketIO) error {
599+
clientIO.ResetSequence()
600+
request, err := clientIO.ReadPacket()
601+
require.NoError(ts.t, err)
602+
rsErr := ts.mp.ExecuteCmd(context.Background(), request)
603+
require.True(ts.t, pnet.IsMySQLError(rsErr))
604+
require.Equal(ts.t, SrcNone, ts.mp.QuitSource())
605+
return nil
606+
},
607+
backend: func(packetIO pnet.PacketIO) error {
608+
ts.mb.respondType = responseTypeErr
609+
return ts.mb.respond(packetIO)
610+
},
611+
},
612+
}
613+
ts.runTests(runners)
614+
}
615+
578616
// Test that closing the BackendConnMgr while it's receiving a redirection signal is OK.
579617
func TestCloseWhileRedirect(t *testing.T) {
580618
ts := newBackendMgrTester(t)

pkg/proxy/client/client_conn.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -75,7 +75,7 @@ func (cc *ClientConnection) processMsg(ctx context.Context) error {
7575
return err
7676
}
7777
err = cc.connMgr.ExecuteCmd(ctx, clientPkt)
78-
if err != nil {
78+
if err != nil && !pnet.IsMySQLError(err) {
7979
return err
8080
}
8181
if pnet.Command(clientPkt[0]) == pnet.ComQuit {

pkg/proxy/net/packetio.go

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -366,25 +366,28 @@ func (p *packetIO) WritePacket(data []byte, flush bool) (err error) {
366366
func (p *packetIO) ForwardUntil(destIO PacketIO, isEnd func(firstByte byte, firstPktLen int) (end, needData bool),
367367
process func(response []byte) error) error {
368368
p.readWriter.BeginRW(rwRead)
369-
dest := destIO.(*packetIO)
370-
dest.readWriter.BeginRW(rwWrite)
369+
dest, _ := destIO.(*packetIO)
370+
// destIO is not packetIO in traffic replay.
371+
if dest != nil {
372+
dest.readWriter.BeginRW(rwWrite)
373+
}
371374
p.limitReader.R = p.readWriter
372375
for {
373376
header, err := p.readWriter.Peek(5)
374377
if err != nil {
375-
return p.wrapErr(errors.Wrap(err, ErrReadConn))
378+
return p.wrapErr(errors.Wrap(errors.WithStack(err), ErrReadConn))
376379
}
377380
length := int(header[0]) | int(header[1])<<8 | int(header[2])<<16
378381
end, needData := isEnd(header[4], length)
379382
var data []byte
380383
// Just call ReadFrom if the caller doesn't need the data, even if it's the last packet.
381-
if end && needData {
384+
if (end && needData) || dest == nil {
382385
// TODO: allocate a buffer from pool and return the buffer after `process`.
383386
data, err = p.ReadPacket()
384387
if err != nil {
385388
return p.wrapErr(errors.Wrap(err, ErrReadConn))
386389
}
387-
if err := dest.WritePacket(data, false); err != nil {
390+
if err := destIO.WritePacket(data, false); err != nil {
388391
return p.wrapErr(errors.Wrap(err, ErrWriteConn))
389392
}
390393
} else {

pkg/proxy/proxy.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -150,8 +150,8 @@ func (s *SQLServer) onConn(ctx context.Context, conn net.Conn, addr string) {
150150
return false, nil, 0, nil
151151
}
152152

153-
connID := s.mu.connID
154153
s.mu.connID++
154+
connID := s.mu.connID
155155
logger := s.logger.With(zap.Uint64("connID", connID), zap.String("client_addr", conn.RemoteAddr().String()),
156156
zap.String("addr", addr))
157157
clientConn := client.NewClientConnection(logger.Named("conn"), conn, s.certMgr.ServerSQLTLS(), s.certMgr.SQLTLS(),

pkg/proxy/proxy_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -258,7 +258,7 @@ func TestRecoverPanic(t *testing.T) {
258258
require.NoError(t, err)
259259
server, err := NewSQLServer(lg, &config.Config{}, certManager, nil, &mockHsHandler{
260260
handshakeResp: func(ctx backend.ConnContext, _ *pnet.HandshakeResp) error {
261-
if ctx.Value(backend.ConnContextKeyConnID).(uint64) == 0 {
261+
if ctx.Value(backend.ConnContextKeyConnID).(uint64) == 1 {
262262
panic("HandleHandshakeResp panic")
263263
}
264264
return nil

pkg/sqlreplay/capture/capture.go

Lines changed: 39 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -86,13 +86,17 @@ var _ Capture = (*capture)(nil)
8686

8787
type capture struct {
8888
sync.Mutex
89-
cfg CaptureConfig
90-
wg waitgroup.WaitGroup
91-
cancel context.CancelFunc
92-
cmdCh chan *cmd.Command
93-
err error
94-
startTime time.Time
95-
lg *zap.Logger
89+
cfg CaptureConfig
90+
wg waitgroup.WaitGroup
91+
cancel context.CancelFunc
92+
cmdCh chan *cmd.Command
93+
err error
94+
startTime time.Time
95+
endTime time.Time
96+
progress float64
97+
capturedCmds uint64
98+
filteredCmds uint64
99+
lg *zap.Logger
96100
}
97101

98102
func NewCapture(lg *zap.Logger) *capture {
@@ -115,6 +119,10 @@ func (c *capture) Start(cfg CaptureConfig) error {
115119
c.stopNoLock(nil)
116120
c.cfg = cfg
117121
c.startTime = time.Now()
122+
c.endTime = time.Time{}
123+
c.progress = 0
124+
c.capturedCmds = 0
125+
c.filteredCmds = 0
118126
c.err = nil
119127
childCtx, cancel := context.WithTimeout(context.Background(), c.cfg.Duration)
120128
c.cancel = cancel
@@ -143,6 +151,7 @@ func (c *capture) collectCmds(bufCh chan<- *bytes.Buffer) {
143151
c.Stop(errors.Wrapf(err, "failed to encode command"))
144152
return
145153
}
154+
c.capturedCmds++
146155
if buf.Len() > c.cfg.flushThreshold {
147156
select {
148157
case bufCh <- buf:
@@ -208,7 +217,7 @@ func (c *capture) Progress() (float64, error) {
208217
c.Lock()
209218
defer c.Unlock()
210219
if c.startTime.IsZero() || c.cfg.Duration == 0 {
211-
return 0, c.err
220+
return c.progress, c.err
212221
}
213222
return float64(time.Since(c.startTime)) / float64(c.cfg.Duration), c.err
214223
}
@@ -219,16 +228,33 @@ func (c *capture) stopNoLock(err error) {
219228
if c.startTime.IsZero() {
220229
return
221230
}
222-
if err != nil {
223-
c.lg.Error("stop capture", zap.Error(err))
224-
}
225-
c.err = err
226231
if c.cancel != nil {
227232
c.cancel()
228233
c.cancel = nil
229234
}
230-
c.startTime = time.Time{}
231235
close(c.cmdCh)
236+
237+
c.endTime = time.Now()
238+
fields := []zap.Field{
239+
zap.Time("start_time", c.startTime),
240+
zap.Time("end_time", c.endTime),
241+
zap.Uint64("captured_cmds", c.capturedCmds),
242+
}
243+
if err != nil {
244+
if c.cfg.Duration > 0 {
245+
c.progress = float64(c.endTime.Sub(c.startTime)) / float64(c.cfg.Duration)
246+
if c.progress > 1 {
247+
c.progress = 1
248+
}
249+
}
250+
c.err = err
251+
fields = append(fields, zap.Error(err))
252+
c.lg.Error("capture failed", fields...)
253+
} else {
254+
c.progress = 1
255+
c.lg.Info("capture finished", fields...)
256+
}
257+
c.startTime = time.Time{}
232258
}
233259

234260
func (c *capture) Stop(err error) {

pkg/sqlreplay/capture/capture_test.go

Lines changed: 39 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -11,16 +11,15 @@ import (
1111
"time"
1212

1313
"github.com/pingcap/tiproxy/lib/util/errors"
14-
"github.com/pingcap/tiproxy/lib/util/logger"
1514
"github.com/pingcap/tiproxy/lib/util/waitgroup"
1615
pnet "github.com/pingcap/tiproxy/pkg/proxy/net"
1716
"github.com/pingcap/tiproxy/pkg/sqlreplay/store"
1817
"github.com/stretchr/testify/require"
18+
"go.uber.org/zap"
1919
)
2020

2121
func TestStartAndStop(t *testing.T) {
22-
lg, _ := logger.CreateLoggerForTest(t)
23-
cpt := NewCapture(lg)
22+
cpt := NewCapture(zap.NewNop())
2423
defer cpt.Close()
2524

2625
packet := append([]byte{pnet.ComQuery.Byte()}, []byte("select 1")...)
@@ -35,15 +34,12 @@ func TestStartAndStop(t *testing.T) {
3534
// start capture and the traffic should be outputted
3635
require.NoError(t, cpt.Start(cfg))
3736
cpt.Capture(packet, time.Now(), 100)
38-
_, err := cpt.Progress()
39-
require.NoError(t, err)
4037
cpt.Stop(errors.Errorf("mock error"))
41-
_, err = cpt.Progress()
42-
require.ErrorContains(t, err, "mock error")
4338
cpt.wg.Wait()
4439
data := writer.getData()
4540
require.Greater(t, len(data), 0)
4641
require.Contains(t, string(data), "select 1")
42+
require.Equal(t, uint64(1), cpt.capturedCmds)
4743

4844
// stop capture and traffic should not be outputted
4945
cpt.Capture(packet, time.Now(), 100)
@@ -65,8 +61,7 @@ func TestStartAndStop(t *testing.T) {
6561
}
6662

6763
func TestConcurrency(t *testing.T) {
68-
lg, _ := logger.CreateLoggerForTest(t)
69-
cpt := NewCapture(lg)
64+
cpt := NewCapture(zap.NewNop())
7065
defer cpt.Close()
7166

7267
writer := newMockWriter(store.WriterCfg{})
@@ -145,3 +140,38 @@ func TestCaptureCfgError(t *testing.T) {
145140
require.Equal(t, maxBuffers, cfg.maxBuffers)
146141
require.Equal(t, maxPendingCommands, cfg.maxPendingCommands)
147142
}
143+
144+
func TestProgress(t *testing.T) {
145+
cpt := NewCapture(zap.NewNop())
146+
defer cpt.Close()
147+
148+
writer := newMockWriter(store.WriterCfg{})
149+
cfg := CaptureConfig{
150+
Output: t.TempDir(),
151+
Duration: 10 * time.Second,
152+
cmdLogger: writer,
153+
}
154+
setStartTime := func(t time.Time) {
155+
cpt.Lock()
156+
cpt.startTime = t
157+
cpt.Unlock()
158+
}
159+
160+
now := time.Now()
161+
require.NoError(t, cpt.Start(cfg))
162+
progress, err := cpt.Progress()
163+
require.NoError(t, err)
164+
require.Less(t, progress, 0.3)
165+
166+
setStartTime(now.Add(-5 * time.Second))
167+
progress, err = cpt.Progress()
168+
require.NoError(t, err)
169+
require.GreaterOrEqual(t, progress, 0.5)
170+
171+
cpt.Stop(errors.Errorf("mock error"))
172+
cpt.wg.Wait()
173+
progress, err = cpt.Progress()
174+
require.ErrorContains(t, err, "mock error")
175+
require.GreaterOrEqual(t, progress, 0.5)
176+
require.Less(t, progress, 1.0)
177+
}

0 commit comments

Comments
 (0)