首次推送
This commit is contained in:
@@ -0,0 +1,104 @@
|
||||
package transport
|
||||
|
||||
// maxUint32Val is ^uint32(0) stored in a variable so that int(maxUint32Val)
|
||||
// is a runtime conversion (avoids "constant overflows int" on 32-bit).
|
||||
var maxUint32Val = ^uint32(0)
|
||||
|
||||
type GenerationWindow struct {
|
||||
mappedBaseOffset int
|
||||
generation uint32
|
||||
mod int
|
||||
receiveWindow int
|
||||
}
|
||||
|
||||
func NewGenerationWindow(mod int, windowSize int) *GenerationWindow {
|
||||
return &GenerationWindow{
|
||||
mod: mod,
|
||||
receiveWindow: windowSize,
|
||||
}
|
||||
}
|
||||
|
||||
func (g *GenerationWindow) Advance(amount int) {
|
||||
if amount <= 0 {
|
||||
return
|
||||
}
|
||||
newBaseOffset := g.mappedBaseOffset + amount
|
||||
genStep := newBaseOffset / g.mod
|
||||
if genStep > 0 {
|
||||
// Cap genStep so the uint32 cast below is safe.
|
||||
// On 64-bit platforms genStep could exceed uint32 range; use uint64 comparison.
|
||||
// On 32-bit platforms genStep ≤ MaxInt32 < MaxUint32, so this is always false.
|
||||
if uint64(genStep) > uint64(maxUint32Val) {
|
||||
genStep = int(maxUint32Val)
|
||||
}
|
||||
g.generation += uint32(genStep)
|
||||
}
|
||||
g.mappedBaseOffset = newBaseOffset % g.mod
|
||||
}
|
||||
|
||||
func (g *GenerationWindow) AdvanceToExcluded(mappedValue int) {
|
||||
moveDist := mappedValue - g.mappedBaseOffset
|
||||
if moveDist < 0 {
|
||||
moveDist += g.mod
|
||||
}
|
||||
g.Advance(moveDist + 1)
|
||||
}
|
||||
|
||||
// SyncTo advances the window baseline toward mappedValue (handles wrap and resync).
|
||||
func (g *GenerationWindow) SyncTo(mappedValue int) {
|
||||
moveDist := mappedValue - g.mappedBaseOffset
|
||||
if moveDist < 0 {
|
||||
moveDist += g.mod
|
||||
}
|
||||
g.Advance(moveDist)
|
||||
}
|
||||
|
||||
func (g *GenerationWindow) IsInWindow(mappedValue int) bool {
|
||||
maxOffset := g.mappedBaseOffset + g.receiveWindow
|
||||
if maxOffset < g.mod {
|
||||
return mappedValue >= g.mappedBaseOffset && mappedValue < maxOffset
|
||||
}
|
||||
|
||||
return mappedValue >= g.mappedBaseOffset || mappedValue < maxOffset-g.mod
|
||||
}
|
||||
|
||||
// MappedToIndex returns the offset from the window base; negative means stale, >= receiveWindow too far ahead.
|
||||
func (g *GenerationWindow) MappedToIndex(mappedValue int) int {
|
||||
if g.IsNextGen(mappedValue) {
|
||||
return (mappedValue + g.mod) - g.mappedBaseOffset
|
||||
}
|
||||
|
||||
return mappedValue - g.mappedBaseOffset
|
||||
}
|
||||
|
||||
// IsOldPacket reports whether mappedValue is before the receive window.
|
||||
func (g *GenerationWindow) IsOldPacket(mappedValue int) bool {
|
||||
index := g.MappedToIndex(mappedValue)
|
||||
|
||||
return index < 0
|
||||
}
|
||||
|
||||
// IsFuturePacket reports whether mappedValue lies beyond the window.
|
||||
func (g *GenerationWindow) IsFuturePacket(mappedValue int) bool {
|
||||
index := g.MappedToIndex(mappedValue)
|
||||
|
||||
return index >= g.receiveWindow
|
||||
}
|
||||
|
||||
func (g *GenerationWindow) IsNextGen(mappedValue int) bool {
|
||||
return g.mappedBaseOffset > (g.mod-g.receiveWindow) &&
|
||||
mappedValue < (g.mappedBaseOffset+g.receiveWindow)-g.mod
|
||||
}
|
||||
|
||||
func (g *GenerationWindow) GetGeneration(mappedValue int) uint32 {
|
||||
if g.IsNextGen(mappedValue) {
|
||||
return g.generation + 1
|
||||
}
|
||||
|
||||
return g.generation
|
||||
}
|
||||
|
||||
func (g *GenerationWindow) Reset() {
|
||||
g.mappedBaseOffset = 0
|
||||
g.generation = 0
|
||||
}
|
||||
@@ -0,0 +1,201 @@
|
||||
package transport_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/honeybbq/teamspeak-go/transport"
|
||||
)
|
||||
|
||||
func TestGenerationWindowNew(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
if gw == nil {
|
||||
t.Fatal("expected non-nil GenerationWindow")
|
||||
}
|
||||
if gw.GetGeneration(0) != 0 {
|
||||
t.Error("initial generation should be 0")
|
||||
}
|
||||
if !gw.IsInWindow(0) {
|
||||
t.Error("0 should be in window after creation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowReset(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
gw.Advance(500)
|
||||
gw.Reset()
|
||||
if gw.GetGeneration(0) != 0 {
|
||||
t.Error("generation should be 0 after Reset")
|
||||
}
|
||||
// After Reset, base=0 so window is [0..1023]
|
||||
if !gw.IsInWindow(0) {
|
||||
t.Error("0 should be in window after Reset")
|
||||
}
|
||||
if !gw.IsInWindow(1023) {
|
||||
t.Error("1023 should be in window after Reset")
|
||||
}
|
||||
if gw.IsInWindow(1024) {
|
||||
t.Error("1024 should NOT be in window after Reset")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowAdvanceBasic(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
gw.Advance(100)
|
||||
if !gw.IsInWindow(100) {
|
||||
t.Error("100 should be in window after Advance(100)")
|
||||
}
|
||||
if gw.IsInWindow(99) {
|
||||
t.Error("99 should not be in window after Advance(100)")
|
||||
}
|
||||
if !gw.IsInWindow(1123) {
|
||||
t.Error("1123 (100+1023) should be in window")
|
||||
}
|
||||
if gw.IsInWindow(1124) {
|
||||
t.Error("1124 should not be in window")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowAdvanceZeroAndNegative(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
gw.Advance(0)
|
||||
gw.Advance(-1)
|
||||
if !gw.IsInWindow(0) {
|
||||
t.Error("0 should still be in window after zero/negative advance")
|
||||
}
|
||||
if gw.GetGeneration(0) != 0 {
|
||||
t.Error("generation should remain 0")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowFullCycleIncrementsGeneration(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
gw.Advance(65536)
|
||||
if gw.GetGeneration(0) != 1 {
|
||||
t.Errorf("after one full cycle, generation should be 1, got %d", gw.GetGeneration(0))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowIsOldPacket(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
gw.Advance(100)
|
||||
if !gw.IsOldPacket(0) {
|
||||
t.Error("0 should be old after Advance(100)")
|
||||
}
|
||||
if !gw.IsOldPacket(99) {
|
||||
t.Error("99 should be old after Advance(100)")
|
||||
}
|
||||
if gw.IsOldPacket(100) {
|
||||
t.Error("100 should not be old (it's the window start)")
|
||||
}
|
||||
if gw.IsOldPacket(500) {
|
||||
t.Error("500 should not be old (it's in window)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowIsFuturePacket(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
if !gw.IsFuturePacket(1024) {
|
||||
t.Error("1024 should be a future packet (base=0, window=1024)")
|
||||
}
|
||||
if gw.IsFuturePacket(1023) {
|
||||
t.Error("1023 should not be a future packet (last in window)")
|
||||
}
|
||||
if gw.IsFuturePacket(500) {
|
||||
t.Error("500 should not be a future packet")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowAdvanceToExcluded(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
// AdvanceToExcluded(10): moveDist=10, Advance(11), base becomes 11
|
||||
gw.AdvanceToExcluded(10)
|
||||
if !gw.IsOldPacket(10) {
|
||||
t.Error("10 should be old after AdvanceToExcluded(10)")
|
||||
}
|
||||
if gw.IsOldPacket(11) {
|
||||
t.Error("11 should be in window (new base)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowSyncTo(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
gw.SyncTo(500)
|
||||
if gw.IsOldPacket(500) {
|
||||
t.Error("500 should not be old after SyncTo(500)")
|
||||
}
|
||||
if !gw.IsOldPacket(499) {
|
||||
t.Error("499 should be old after SyncTo(500)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowWrapAroundInWindow(t *testing.T) {
|
||||
// mod=16, window=4: after Advance(14), window spans 14,15,0,1
|
||||
gw := transport.NewGenerationWindow(16, 4)
|
||||
gw.Advance(14)
|
||||
if !gw.IsInWindow(14) {
|
||||
t.Error("14 should be in window")
|
||||
}
|
||||
if !gw.IsInWindow(15) {
|
||||
t.Error("15 should be in window")
|
||||
}
|
||||
if !gw.IsInWindow(0) {
|
||||
t.Error("0 (wrapped) should be in window")
|
||||
}
|
||||
if !gw.IsInWindow(1) {
|
||||
t.Error("1 (wrapped) should be in window")
|
||||
}
|
||||
if gw.IsInWindow(2) {
|
||||
t.Error("2 should NOT be in window")
|
||||
}
|
||||
if gw.IsInWindow(13) {
|
||||
t.Error("13 should NOT be in window")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowIsNextGen(t *testing.T) {
|
||||
// mod=16, window=4: IsNextGen requires base > 16-4=12
|
||||
gw := transport.NewGenerationWindow(16, 4)
|
||||
gw.Advance(13) // base=13
|
||||
// 0 < (13+4)-16=1 and base=13>12 → IsNextGen(0)=true
|
||||
if !gw.IsNextGen(0) {
|
||||
t.Error("0 should be next-gen when base=13, mod=16, window=4")
|
||||
}
|
||||
// 1 is NOT < 1 → IsNextGen(1)=false
|
||||
if gw.IsNextGen(1) {
|
||||
t.Error("1 should NOT be next-gen when base=13")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowNextGenIncreasesGeneration(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(16, 4)
|
||||
gw.Advance(13)
|
||||
if gw.GetGeneration(13) != 0 {
|
||||
t.Errorf("generation of 13 should be 0, got %d", gw.GetGeneration(13))
|
||||
}
|
||||
if gw.GetGeneration(0) != 1 {
|
||||
t.Errorf("generation of 0 (next-gen) should be 1, got %d", gw.GetGeneration(0))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowMultipleAdvanceCycles(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
gw.Advance(65536)
|
||||
gw.Advance(65536)
|
||||
gw.Advance(65536)
|
||||
if gw.GetGeneration(0) != 3 {
|
||||
t.Errorf("after 3 full cycles, generation should be 3, got %d", gw.GetGeneration(0))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerationWindowMappedToIndex(t *testing.T) {
|
||||
gw := transport.NewGenerationWindow(65536, 1024)
|
||||
gw.Advance(100) // base=100
|
||||
idx := gw.MappedToIndex(150)
|
||||
if idx != 50 {
|
||||
t.Errorf("MappedToIndex(150)=%d, want 50 (base=100)", idx)
|
||||
}
|
||||
idx = gw.MappedToIndex(99)
|
||||
if idx != -1 {
|
||||
t.Errorf("MappedToIndex(99)=%d, want -1 (old packet)", idx)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,945 @@
|
||||
package transport
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"github.com/honeybbq/teamspeak-go/crypto"
|
||||
"github.com/honeybbq/teamspeak-go/handshake"
|
||||
)
|
||||
|
||||
const (
|
||||
MaxOutPacketSize = 500
|
||||
ReceivePacketWindowSize = 1024
|
||||
PingInterval = 5 * time.Second
|
||||
PacketTimeout = 60 * time.Second
|
||||
MaxRetryInterval = time.Second
|
||||
udpReadBufferSize = 4096
|
||||
voicePayloadBufferSize = 1027
|
||||
smallPacketBufferSize = 2
|
||||
packetProcessQueueSize = 2048
|
||||
resendBaseInterval = 500 * time.Millisecond
|
||||
resendLoopInterval = 100 * time.Millisecond
|
||||
headerSize = 5
|
||||
tagSize = 8
|
||||
voiceHeaderSize = 3
|
||||
ackDataSize = 2
|
||||
)
|
||||
|
||||
var bufPool = sync.Pool{
|
||||
New: func() any {
|
||||
buf := make([]byte, udpReadBufferSize)
|
||||
|
||||
return &buf
|
||||
},
|
||||
}
|
||||
|
||||
// voicePayloadPool holds buffers for the 3-byte voice header plus Opus frame.
|
||||
var voicePayloadPool = sync.Pool{
|
||||
New: func() any {
|
||||
buf := make([]byte, voicePayloadBufferSize)
|
||||
|
||||
return &buf
|
||||
},
|
||||
}
|
||||
|
||||
// smallBufPool holds 2-byte buffers for ACK/Pong payloads.
|
||||
var smallBufPool = sync.Pool{
|
||||
New: func() any {
|
||||
buf := make([]byte, smallPacketBufferSize)
|
||||
|
||||
return &buf
|
||||
},
|
||||
}
|
||||
|
||||
type pooledBuffer struct {
|
||||
buf []byte
|
||||
n int
|
||||
}
|
||||
|
||||
type PacketHandler struct {
|
||||
lastMessageReceived time.Time
|
||||
conn io.ReadWriteCloser
|
||||
commandQueue map[uint16]*Packet
|
||||
commandLowQueue map[uint16]*Packet
|
||||
stopCh chan struct{}
|
||||
recvWindowCommand *GenerationWindow
|
||||
recvWindowCommandLow *GenerationWindow
|
||||
sendWindowCommand *GenerationWindow
|
||||
sendWindowCommandLow *GenerationWindow
|
||||
ackManager map[uint32]*resendPacket
|
||||
initPacketCheck *resendPacket
|
||||
packetProcessCh chan *pooledBuffer
|
||||
OnClosed func(err error)
|
||||
logger *slog.Logger
|
||||
OnAck func(id uint16)
|
||||
OnPacket func(p *Packet)
|
||||
TsCrypt *crypto.Crypt
|
||||
generationCounter [9]uint32
|
||||
mu sync.Mutex
|
||||
closed atomic.Bool
|
||||
packetCounter [9]uint16
|
||||
clientID uint16
|
||||
nextCommandLowID uint16
|
||||
nextCommandID uint16
|
||||
}
|
||||
|
||||
type resendPacket struct {
|
||||
packet *Packet
|
||||
firstSend time.Time
|
||||
lastSend time.Time
|
||||
retryCount int
|
||||
nextInterval time.Duration
|
||||
}
|
||||
|
||||
type decryptPacketResult struct {
|
||||
plaintext []byte
|
||||
dummyUsed bool
|
||||
}
|
||||
|
||||
func NewPacketHandler(tsCrypt *crypto.Crypt, logger *slog.Logger) *PacketHandler {
|
||||
if logger == nil {
|
||||
logger = slog.Default()
|
||||
}
|
||||
|
||||
return &PacketHandler{
|
||||
TsCrypt: tsCrypt,
|
||||
logger: logger,
|
||||
ackManager: make(map[uint32]*resendPacket),
|
||||
packetProcessCh: make(chan *pooledBuffer, packetProcessQueueSize),
|
||||
stopCh: make(chan struct{}),
|
||||
recvWindowCommand: NewGenerationWindow(1<<16, ReceivePacketWindowSize),
|
||||
recvWindowCommandLow: NewGenerationWindow(1<<16, ReceivePacketWindowSize),
|
||||
sendWindowCommand: NewGenerationWindow(1<<16, ReceivePacketWindowSize),
|
||||
sendWindowCommandLow: NewGenerationWindow(1<<16, ReceivePacketWindowSize),
|
||||
commandQueue: make(map[uint16]*Packet),
|
||||
commandLowQueue: make(map[uint16]*Packet),
|
||||
lastMessageReceived: time.Now(),
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) SetClientID(id uint16) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
h.clientID = id
|
||||
}
|
||||
|
||||
// Connect resolves addr as a UDP address, dials it, and calls Start.
|
||||
func (h *PacketHandler) Connect(addr string) error {
|
||||
udpAddr, err := net.ResolveUDPAddr("udp", addr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
conn, err := net.DialUDP("udp", nil, udpAddr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
return h.Start(conn)
|
||||
}
|
||||
|
||||
// Start attaches conn to the handler, launches background goroutines, and sends
|
||||
// the initial Init1 handshake packet. conn must implement io.ReadWriteCloser
|
||||
// with datagram-style semantics (each Write produces one discrete message).
|
||||
func (h *PacketHandler) Start(conn io.ReadWriteCloser) error {
|
||||
h.conn = conn
|
||||
|
||||
go h.receiveLoop()
|
||||
go h.processLoop()
|
||||
go h.resendLoop()
|
||||
go h.pingLoop()
|
||||
|
||||
h.packetCounter[PacketTypeCommand]++
|
||||
h.packetCounter[PacketTypeInit1] = 101
|
||||
|
||||
init1Data := handshake.ProcessInit1(h.TsCrypt, nil)
|
||||
|
||||
return h.SendPacket(byte(PacketTypeInit1), init1Data, 0)
|
||||
}
|
||||
|
||||
func (h *PacketHandler) SendPacket(pType byte, data []byte, flags byte) error {
|
||||
dummy := !h.TsCrypt.CryptoInitComplete
|
||||
|
||||
// Fragment non-voice command payloads larger than one UDP frame (487 B body).
|
||||
if len(data) > 487 && pType != 0 && pType != 1 {
|
||||
return h.sendSplitPacket(pType, data, flags, dummy)
|
||||
}
|
||||
|
||||
return h.sendPacket(pType, data, flags, dummy)
|
||||
}
|
||||
|
||||
func (h *PacketHandler) sendSplitPacket(pType byte, data []byte, flags byte, dummy bool) error {
|
||||
maxSize := 487 // MaxOutPacketSize(500) - Header(5) - Tag(8)
|
||||
pos := 0
|
||||
first := true
|
||||
|
||||
for pos < len(data) {
|
||||
blockSize := min(len(data)-pos, maxSize)
|
||||
|
||||
last := (pos + blockSize) == len(data)
|
||||
|
||||
pFlags := flags
|
||||
// TeamSpeak sets Fragmented on the first and last chunk only.
|
||||
if first != last {
|
||||
pFlags |= byte(PacketFlagFragmented)
|
||||
}
|
||||
|
||||
err := h.sendPacket(pType, data[pos:pos+blockSize], pFlags, dummy)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
pos += blockSize
|
||||
first = false
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (h *PacketHandler) sendPacket(pType byte, data []byte, flags byte, dummy bool) error {
|
||||
flags = applyProtocolFlags(pType, flags)
|
||||
pID, pGen := h.nextPacketIdentity(pType)
|
||||
|
||||
p := &Packet{
|
||||
TypeFlagged: pType | flags,
|
||||
ID: pID,
|
||||
GenerationID: pGen,
|
||||
Data: data,
|
||||
ClientID: h.clientID,
|
||||
}
|
||||
|
||||
unencrypted := (flags&byte(PacketFlagUnencrypted) != 0)
|
||||
header := p.BuildC2SHeader()
|
||||
ciphertext, tag, err := h.TsCrypt.Encrypt(pType, p.ID, p.GenerationID, header, p.Data, dummy, unencrypted)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
final := getPooledBytes(&bufPool, tagSize+headerSize+len(ciphertext))
|
||||
defer putPooledBytes(&bufPool, final)
|
||||
|
||||
copy(final[0:8], tag)
|
||||
copy(final[8:13], header)
|
||||
copy(final[13:], ciphertext)
|
||||
|
||||
_, err = h.conn.Write(final[:tagSize+headerSize+len(ciphertext)])
|
||||
|
||||
rp := &resendPacket{
|
||||
packet: p,
|
||||
firstSend: time.Now(),
|
||||
lastSend: time.Now(),
|
||||
nextInterval: resendBaseInterval,
|
||||
}
|
||||
h.trackResendPacket(pType, p, rp)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *PacketHandler) sendPong(pID uint16, dummy bool) error {
|
||||
pongData := getPooledBytes(&smallBufPool, smallPacketBufferSize)
|
||||
binary.BigEndian.PutUint16(pongData, pID)
|
||||
err := h.sendPacket(byte(PacketTypePong), pongData, byte(PacketFlagUnencrypted), dummy)
|
||||
putPooledBytes(&smallBufPool, pongData)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *PacketHandler) receiveLoop() {
|
||||
var finalErr error
|
||||
defer func() {
|
||||
if h.OnClosed != nil {
|
||||
h.OnClosed(finalErr)
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
buf := getPooledBytes(&bufPool, udpReadBufferSize)
|
||||
n, err := h.conn.Read(buf)
|
||||
if err != nil {
|
||||
putPooledBytes(&bufPool, buf)
|
||||
select {
|
||||
case <-h.stopCh:
|
||||
return
|
||||
default:
|
||||
h.logger.Error("udp read failed", slog.Any("error", err))
|
||||
finalErr = err
|
||||
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
select {
|
||||
case h.packetProcessCh <- &pooledBuffer{buf: buf, n: n}:
|
||||
default:
|
||||
h.logger.Warn("packet process channel full, dropping packet")
|
||||
putPooledBytes(&bufPool, buf)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) processLoop() {
|
||||
for {
|
||||
select {
|
||||
case <-h.stopCh:
|
||||
return
|
||||
case pb := <-h.packetProcessCh:
|
||||
h.handleRawPacket(pb.buf[:pb.n])
|
||||
putPooledBytes(&bufPool, pb.buf)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) handleRawPacket(raw []byte) {
|
||||
if len(raw) < 11 {
|
||||
return
|
||||
}
|
||||
|
||||
tag := raw[0:8]
|
||||
header := raw[8:11]
|
||||
ciphertext := raw[11:]
|
||||
p := parseServerPacket(header)
|
||||
p.ReceivedAt = time.Now()
|
||||
h.markMessageReceived()
|
||||
|
||||
decrypted, ok := h.decryptPacketData(p, header, ciphertext, tag)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
p.Data = decrypted.plaintext
|
||||
|
||||
if p.Type() == PacketTypePing {
|
||||
_ = h.sendPong(p.ID, decrypted.dummyUsed)
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
if !h.handleCommandWindowAndAck(p, decrypted.dummyUsed) {
|
||||
return
|
||||
}
|
||||
|
||||
h.handlePacketQueue(p)
|
||||
h.updatePostReceiveState(p)
|
||||
}
|
||||
|
||||
func (h *PacketHandler) getWinForType(pType PacketType) *GenerationWindow {
|
||||
switch pType {
|
||||
case PacketTypeCommand:
|
||||
return h.recvWindowCommand
|
||||
case PacketTypeCommandLow:
|
||||
return h.recvWindowCommandLow
|
||||
case PacketTypeVoice, PacketTypeVoiceWhisper, PacketTypePing, PacketTypePong,
|
||||
PacketTypeAck, PacketTypeAckLow, PacketTypeInit1:
|
||||
return nil
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) handlePacketQueue(p *Packet) {
|
||||
pType := p.Type()
|
||||
if pType != PacketTypeCommand && pType != PacketTypeCommandLow {
|
||||
if h.OnPacket != nil {
|
||||
h.OnPacket(p)
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
h.mu.Lock()
|
||||
var queue map[uint16]*Packet
|
||||
var nextID *uint16
|
||||
if pType == PacketTypeCommand {
|
||||
queue = h.commandQueue
|
||||
nextID = &h.nextCommandID
|
||||
} else {
|
||||
queue = h.commandLowQueue
|
||||
nextID = &h.nextCommandLowID
|
||||
}
|
||||
|
||||
queue[p.ID] = p
|
||||
|
||||
// If the expected ID never arrives, skip it once a newer fragment has stalled long enough.
|
||||
h.fastForwardMissingPackets(pType, queue, nextID)
|
||||
|
||||
for {
|
||||
packet, ok := queue[*nextID]
|
||||
if !ok {
|
||||
h.logQueueBacklog(pType, queue, *nextID)
|
||||
|
||||
break
|
||||
}
|
||||
|
||||
var win *GenerationWindow
|
||||
if pType == PacketTypeCommand {
|
||||
win = h.recvWindowCommand
|
||||
} else {
|
||||
win = h.recvWindowCommandLow
|
||||
}
|
||||
|
||||
reassembled, complete := h.tryReassemble(packet, queue, nextID, win)
|
||||
if !complete {
|
||||
break
|
||||
}
|
||||
|
||||
h.tryDecompressPacket(reassembled)
|
||||
|
||||
if h.OnPacket != nil {
|
||||
h.mu.Unlock()
|
||||
h.OnPacket(reassembled)
|
||||
h.mu.Lock()
|
||||
}
|
||||
}
|
||||
h.mu.Unlock()
|
||||
}
|
||||
|
||||
func parseServerPacket(header []byte) *Packet {
|
||||
p := &Packet{}
|
||||
p.ParseS2CHeader(header)
|
||||
|
||||
return p
|
||||
}
|
||||
|
||||
func (h *PacketHandler) markMessageReceived() {
|
||||
h.mu.Lock()
|
||||
h.lastMessageReceived = time.Now()
|
||||
h.mu.Unlock()
|
||||
}
|
||||
|
||||
func (h *PacketHandler) resolvePacketGeneration(p *Packet) uint32 {
|
||||
var gen uint32
|
||||
h.mu.Lock()
|
||||
switch p.Type() {
|
||||
case PacketTypeCommand:
|
||||
gen = h.recvWindowCommand.GetGeneration(int(p.ID))
|
||||
case PacketTypeCommandLow:
|
||||
gen = h.recvWindowCommandLow.GetGeneration(int(p.ID))
|
||||
case PacketTypeAck:
|
||||
gen = h.sendWindowCommand.GetGeneration(int(p.ID))
|
||||
case PacketTypeAckLow:
|
||||
gen = h.sendWindowCommandLow.GetGeneration(int(p.ID))
|
||||
case PacketTypeVoice, PacketTypeVoiceWhisper, PacketTypePing, PacketTypePong, PacketTypeInit1:
|
||||
// No generation tracking for these packet types.
|
||||
default:
|
||||
// Unknown packet type, keep generation as zero.
|
||||
}
|
||||
h.mu.Unlock()
|
||||
|
||||
return gen
|
||||
}
|
||||
|
||||
func (h *PacketHandler) decryptPacketData(
|
||||
p *Packet, header, ciphertext, tag []byte,
|
||||
) (*decryptPacketResult, bool) {
|
||||
unencrypted := (p.Flags() & PacketFlagUnencrypted) != 0
|
||||
dummy := !h.TsCrypt.CryptoInitComplete
|
||||
dummyUsed := dummy
|
||||
gen := h.resolvePacketGeneration(p)
|
||||
|
||||
plaintext, err := h.TsCrypt.Decrypt(byte(p.Type()), p.ID, gen, header, ciphertext, tag, dummy, unencrypted)
|
||||
if err != nil && !dummy && !unencrypted {
|
||||
plaintext, gen, err = h.decryptWithGenerationGuess(p, gen, header, ciphertext, tag)
|
||||
}
|
||||
if err != nil && !dummy {
|
||||
plaintext, dummyUsed, err = h.decryptWithDummyFallback(
|
||||
p, gen, header, ciphertext, tag, unencrypted, plaintext, dummyUsed, err,
|
||||
)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
return &decryptPacketResult{plaintext: plaintext, dummyUsed: dummyUsed}, true
|
||||
}
|
||||
|
||||
func (h *PacketHandler) decryptWithGenerationGuess(
|
||||
p *Packet, gen uint32, header, ciphertext, tag []byte,
|
||||
) ([]byte, uint32, error) {
|
||||
for _, offset := range []int{-1, 1} {
|
||||
guessGen, ok := shiftGeneration(gen, offset)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
plaintext, err := h.TsCrypt.Decrypt(byte(p.Type()), p.ID, guessGen, header, ciphertext, tag, false, false)
|
||||
if err == nil {
|
||||
h.logger.Debug("generation guess succeeded",
|
||||
slog.Uint64("id", uint64(p.ID)),
|
||||
slog.Int("offset", offset),
|
||||
slog.Uint64("new_gen", uint64(guessGen)))
|
||||
|
||||
return plaintext, guessGen, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, gen, errDecryptFailed
|
||||
}
|
||||
|
||||
var errDecryptFailed = errors.New("decrypt failed")
|
||||
|
||||
func (h *PacketHandler) decryptWithDummyFallback(
|
||||
p *Packet, gen uint32, header, ciphertext, tag []byte, unencrypted bool,
|
||||
plaintext []byte, dummyUsed bool, decryptErr error,
|
||||
) ([]byte, bool, error) {
|
||||
switch p.Type() {
|
||||
case PacketTypeCommand, PacketTypeCommandLow, PacketTypeAck:
|
||||
plaintext, decryptErr = h.TsCrypt.Decrypt(byte(p.Type()), p.ID, gen, header, ciphertext, tag, true, unencrypted)
|
||||
if decryptErr == nil {
|
||||
return plaintext, true, nil
|
||||
}
|
||||
case PacketTypeVoice, PacketTypeVoiceWhisper, PacketTypePing, PacketTypePong, PacketTypeAckLow, PacketTypeInit1:
|
||||
// No dummy fallback path required.
|
||||
default:
|
||||
// Unknown packet type.
|
||||
}
|
||||
h.logger.Debug("packet decryption failed",
|
||||
slog.Uint64("type", uint64(p.Type())),
|
||||
slog.Uint64("id", uint64(p.ID)),
|
||||
slog.Uint64("gen", uint64(gen)),
|
||||
slog.Any("error", decryptErr))
|
||||
|
||||
return plaintext, dummyUsed, decryptErr
|
||||
}
|
||||
|
||||
func (h *PacketHandler) handleCommandWindowAndAck(p *Packet, dummyUsed bool) bool {
|
||||
if p.Type() != PacketTypeCommand && p.Type() != PacketTypeCommandLow {
|
||||
return true
|
||||
}
|
||||
h.mu.Lock()
|
||||
var win *GenerationWindow
|
||||
if p.Type() == PacketTypeCommand {
|
||||
win = h.recvWindowCommand
|
||||
} else {
|
||||
win = h.recvWindowCommandLow
|
||||
}
|
||||
inWindow := win.IsInWindow(int(p.ID))
|
||||
isOld := win.IsOldPacket(int(p.ID))
|
||||
h.mu.Unlock()
|
||||
|
||||
ackType := PacketTypeAck
|
||||
if p.Type() == PacketTypeCommandLow {
|
||||
ackType = PacketTypeAckLow
|
||||
}
|
||||
|
||||
if !inWindow {
|
||||
if isOld {
|
||||
h.logger.Debug("received old packet, sending ack only",
|
||||
slog.Uint64("type", uint64(p.Type())),
|
||||
slog.Uint64("id", uint64(p.ID)))
|
||||
h.sendAckPacket(p.ID, ackType, dummyUsed)
|
||||
} else {
|
||||
h.logger.Warn("packet too far ahead, ignoring",
|
||||
slog.Uint64("type", uint64(p.Type())),
|
||||
slog.Uint64("id", uint64(p.ID)))
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
h.logger.Debug("sending ack for command",
|
||||
slog.Uint64("type", uint64(ackType)),
|
||||
slog.Uint64("id", uint64(p.ID)))
|
||||
h.sendAckPacket(p.ID, ackType, dummyUsed)
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func (h *PacketHandler) sendAckPacket(packetID uint16, ackType PacketType, dummyUsed bool) {
|
||||
ackData := getPooledBytes(&smallBufPool, ackDataSize)
|
||||
binary.BigEndian.PutUint16(ackData, packetID)
|
||||
_ = h.sendPacket(byte(ackType), ackData, 0, dummyUsed)
|
||||
putPooledBytes(&smallBufPool, ackData)
|
||||
}
|
||||
|
||||
func (h *PacketHandler) updatePostReceiveState(p *Packet) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
|
||||
if p.Type() == PacketTypeInit1 {
|
||||
h.logger.Debug("received init1 response, cleared init packet check")
|
||||
h.initPacketCheck = nil
|
||||
|
||||
return
|
||||
}
|
||||
if (p.Type() == PacketTypeAck || p.Type() == PacketTypeAckLow) && len(p.Data) >= 2 {
|
||||
ackID := binary.BigEndian.Uint16(p.Data[0:2])
|
||||
targetType := uint32(PacketTypeCommand)
|
||||
if p.Type() == PacketTypeAckLow {
|
||||
targetType = uint32(PacketTypeCommandLow)
|
||||
}
|
||||
h.logger.Debug("received ack from server",
|
||||
slog.Uint64("target_type", uint64(targetType)),
|
||||
slog.Uint64("id", uint64(ackID)))
|
||||
delete(h.ackManager, (targetType<<16)|uint32(ackID))
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) fastForwardMissingPackets(
|
||||
pType PacketType, queue map[uint16]*Packet, nextID *uint16,
|
||||
) {
|
||||
for {
|
||||
if _, ok := queue[*nextID]; ok {
|
||||
return
|
||||
}
|
||||
if !hasOldNewerPacket(queue, *nextID) {
|
||||
return
|
||||
}
|
||||
h.logger.Warn("skipping missing packet to unblock queue",
|
||||
slog.Uint64("type", uint64(pType)),
|
||||
slog.Uint64("missing_id", uint64(*nextID)))
|
||||
*nextID++
|
||||
if win := h.getWinForType(pType); win != nil {
|
||||
win.Advance(1)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func hasOldNewerPacket(queue map[uint16]*Packet, nextID uint16) bool {
|
||||
for id, pkg := range queue {
|
||||
if (id-nextID) < 32768 && time.Since(pkg.ReceivedAt) > 5*time.Second {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func (h *PacketHandler) logQueueBacklog(pType PacketType, queue map[uint16]*Packet, nextID uint16) {
|
||||
if len(queue) <= 10 {
|
||||
return
|
||||
}
|
||||
h.logger.Debug("packet queue backlog",
|
||||
slog.Uint64("type", uint64(pType)),
|
||||
slog.Uint64("next_id", uint64(nextID)),
|
||||
slog.Int("backlog_size", len(queue)))
|
||||
}
|
||||
|
||||
func (h *PacketHandler) tryDecompressPacket(packet *Packet) {
|
||||
if (packet.Flags() & PacketFlagCompressed) == 0 {
|
||||
return
|
||||
}
|
||||
qlz := NewQlz()
|
||||
decompressed, err := qlz.Decompress(packet.Data)
|
||||
if err != nil {
|
||||
h.logger.Debug("decompression failed",
|
||||
slog.Uint64("id", uint64(packet.ID)),
|
||||
slog.Any("error", err))
|
||||
|
||||
return
|
||||
}
|
||||
h.logger.Debug("decompressed packet successfully",
|
||||
slog.Uint64("id", uint64(packet.ID)),
|
||||
slog.Int("old_len", len(packet.Data)),
|
||||
slog.Int("new_len", len(decompressed)))
|
||||
packet.Data = decompressed
|
||||
packet.TypeFlagged &= ^byte(PacketFlagCompressed)
|
||||
}
|
||||
|
||||
func (h *PacketHandler) tryReassemble(
|
||||
startPacket *Packet, queue map[uint16]*Packet, nextID *uint16, win *GenerationWindow,
|
||||
) (*Packet, bool) {
|
||||
if (startPacket.Flags() & PacketFlagFragmented) == 0 {
|
||||
advanceQueueWindow(queue, nextID, win)
|
||||
|
||||
return startPacket, true
|
||||
}
|
||||
|
||||
fragments, totalSize, ok := collectFragments(queue, *nextID)
|
||||
if !ok {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
h.logger.Debug("reassembling fragmented packet",
|
||||
slog.Uint64("start_id", uint64(*nextID)),
|
||||
slog.Int("fragments", len(fragments)),
|
||||
slog.Int("total_size", totalSize))
|
||||
|
||||
combined := make([]byte, totalSize)
|
||||
pos := 0
|
||||
for i := range fragments {
|
||||
copy(combined[pos:], fragments[i].Data)
|
||||
pos += len(fragments[i].Data)
|
||||
advanceQueueWindow(queue, nextID, win)
|
||||
}
|
||||
|
||||
startPacket.Data = combined
|
||||
startPacket.TypeFlagged &= ^byte(PacketFlagFragmented)
|
||||
|
||||
return startPacket, true
|
||||
}
|
||||
|
||||
func applyProtocolFlags(pType byte, flags byte) byte {
|
||||
if pType == byte(PacketTypeCommand) || pType == byte(PacketTypeCommandLow) {
|
||||
return flags | byte(PacketFlagNewProtocol)
|
||||
}
|
||||
|
||||
return flags
|
||||
}
|
||||
|
||||
func (h *PacketHandler) nextPacketIdentity(pType byte) (uint16, uint32) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
|
||||
pID := h.packetCounter[pType]
|
||||
pGen := h.generationCounter[pType]
|
||||
if pType == byte(PacketTypeInit1) {
|
||||
return pID, pGen
|
||||
}
|
||||
|
||||
h.packetCounter[pType]++
|
||||
if h.packetCounter[pType] == 0 {
|
||||
h.generationCounter[pType]++
|
||||
}
|
||||
if pType == byte(PacketTypeCommand) {
|
||||
h.sendWindowCommand.AdvanceToExcluded(int(pID))
|
||||
} else if pType == byte(PacketTypeCommandLow) {
|
||||
h.sendWindowCommandLow.AdvanceToExcluded(int(pID))
|
||||
}
|
||||
|
||||
return pID, pGen
|
||||
}
|
||||
|
||||
func (h *PacketHandler) trackResendPacket(pType byte, p *Packet, rp *resendPacket) {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
|
||||
if pType == byte(PacketTypeInit1) {
|
||||
h.initPacketCheck = rp
|
||||
|
||||
return
|
||||
}
|
||||
if pType == byte(PacketTypeCommand) || pType == byte(PacketTypeCommandLow) {
|
||||
key := (uint32(pType) << 16) | uint32(p.ID)
|
||||
h.ackManager[key] = rp
|
||||
}
|
||||
}
|
||||
|
||||
func collectFragments(queue map[uint16]*Packet, startID uint16) ([]*Packet, int, bool) {
|
||||
var fragments []*Packet
|
||||
currID := startID
|
||||
totalSize := 0
|
||||
startSeen := false
|
||||
for {
|
||||
p, ok := queue[currID]
|
||||
if !ok {
|
||||
return nil, 0, false
|
||||
}
|
||||
fragments = append(fragments, p)
|
||||
totalSize += len(p.Data)
|
||||
var complete bool
|
||||
startSeen, complete = updateFragmentState(startSeen, p.Flags())
|
||||
if complete {
|
||||
return fragments, totalSize, true
|
||||
}
|
||||
currID++
|
||||
}
|
||||
}
|
||||
|
||||
func updateFragmentState(startSeen bool, flags PacketFlags) (bool, bool) {
|
||||
if (flags & PacketFlagFragmented) != 0 {
|
||||
if !startSeen {
|
||||
return true, false
|
||||
}
|
||||
|
||||
return true, true
|
||||
}
|
||||
if !startSeen {
|
||||
return true, true
|
||||
}
|
||||
|
||||
return startSeen, false
|
||||
}
|
||||
|
||||
func shiftGeneration(gen uint32, offset int) (uint32, bool) {
|
||||
switch offset {
|
||||
case -1:
|
||||
if gen == 0 {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
return gen - 1, true
|
||||
case 1:
|
||||
if gen == ^uint32(0) {
|
||||
return 0, false
|
||||
}
|
||||
|
||||
return gen + 1, true
|
||||
default:
|
||||
return gen, true
|
||||
}
|
||||
}
|
||||
|
||||
func advanceQueueWindow(queue map[uint16]*Packet, nextID *uint16, win *GenerationWindow) {
|
||||
delete(queue, *nextID)
|
||||
*nextID++
|
||||
if win != nil {
|
||||
win.Advance(1)
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) pingLoop() {
|
||||
ticker := time.NewTicker(PingInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-h.stopCh:
|
||||
return
|
||||
case <-ticker.C:
|
||||
if h.TsCrypt.CryptoInitComplete {
|
||||
_ = h.SendPacket(byte(PacketTypePing), []byte{}, byte(PacketFlagUnencrypted))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) resendLoop() {
|
||||
ticker := time.NewTicker(resendLoopInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-h.stopCh:
|
||||
return
|
||||
case <-ticker.C:
|
||||
h.checkResends()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) ReceivedFinalInitAck() {
|
||||
h.mu.Lock()
|
||||
h.initPacketCheck = nil
|
||||
h.mu.Unlock()
|
||||
}
|
||||
|
||||
func (h *PacketHandler) checkResends() {
|
||||
h.mu.Lock()
|
||||
now := time.Now()
|
||||
needClose := false
|
||||
|
||||
if now.Sub(h.lastMessageReceived) > PacketTimeout {
|
||||
h.logger.Warn("idle timeout: no packets received", slog.Duration("timeout", PacketTimeout))
|
||||
needClose = true
|
||||
}
|
||||
|
||||
if h.initPacketCheck != nil {
|
||||
h.doResend(h.initPacketCheck, now)
|
||||
}
|
||||
for key, rp := range h.ackManager {
|
||||
if now.Sub(rp.firstSend) > PacketTimeout {
|
||||
delete(h.ackManager, key)
|
||||
needClose = true
|
||||
|
||||
break
|
||||
}
|
||||
h.doResend(rp, now)
|
||||
}
|
||||
h.mu.Unlock()
|
||||
if needClose {
|
||||
_ = h.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) doResend(rp *resendPacket, now time.Time) {
|
||||
if now.Sub(rp.lastSend) >= rp.nextInterval {
|
||||
rp.lastSend = now
|
||||
rp.retryCount++
|
||||
rp.nextInterval *= 2
|
||||
if rp.nextInterval > MaxRetryInterval {
|
||||
rp.nextInterval = MaxRetryInterval
|
||||
}
|
||||
|
||||
unencrypted := (rp.packet.Flags()&PacketFlagUnencrypted != 0)
|
||||
dummy := !h.TsCrypt.CryptoInitComplete
|
||||
header := rp.packet.BuildC2SHeader()
|
||||
h.logger.Debug("resending packet",
|
||||
slog.Uint64("type", uint64(rp.packet.Type())),
|
||||
slog.Uint64("id", uint64(rp.packet.ID)),
|
||||
slog.Int("retry_count", rp.retryCount),
|
||||
slog.Duration("next_interval", rp.nextInterval))
|
||||
ciphertext, tag, _ := h.TsCrypt.Encrypt(
|
||||
byte(rp.packet.Type()), rp.packet.ID, rp.packet.GenerationID, header, rp.packet.Data, dummy, unencrypted,
|
||||
)
|
||||
final := make([]byte, tagSize+headerSize+len(ciphertext))
|
||||
copy(final[0:8], tag)
|
||||
copy(final[8:13], header)
|
||||
copy(final[13:], ciphertext)
|
||||
_, err := h.conn.Write(final)
|
||||
if err != nil {
|
||||
h.logger.Warn("resend write failed", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (h *PacketHandler) SendVoicePacket(data []byte, codec byte) error {
|
||||
h.mu.Lock()
|
||||
pID := h.packetCounter[PacketTypeVoice]
|
||||
pGen := h.generationCounter[PacketTypeVoice]
|
||||
h.packetCounter[PacketTypeVoice]++
|
||||
if h.packetCounter[PacketTypeVoice] == 0 {
|
||||
h.generationCounter[PacketTypeVoice]++
|
||||
}
|
||||
clid := h.clientID
|
||||
h.mu.Unlock()
|
||||
|
||||
payloadLen := voiceHeaderSize + len(data)
|
||||
|
||||
voicePayload := getPooledBytes(&voicePayloadPool, payloadLen)
|
||||
binary.BigEndian.PutUint16(voicePayload[0:2], pID)
|
||||
voicePayload[2] = codec
|
||||
copy(voicePayload[voiceHeaderSize:], data)
|
||||
|
||||
p := &Packet{
|
||||
TypeFlagged: byte(PacketTypeVoice) | byte(PacketFlagUnencrypted),
|
||||
ID: pID,
|
||||
GenerationID: pGen,
|
||||
Data: voicePayload,
|
||||
ClientID: clid,
|
||||
}
|
||||
|
||||
header := p.BuildC2SHeader()
|
||||
|
||||
final := getPooledBytes(&bufPool, tagSize+headerSize+payloadLen)
|
||||
|
||||
copy(final[0:8], h.TsCrypt.FakeSignature)
|
||||
copy(final[8:13], header)
|
||||
copy(final[13:], voicePayload)
|
||||
|
||||
_, err := h.conn.Write(final[:tagSize+headerSize+payloadLen])
|
||||
|
||||
putPooledBytes(&bufPool, final)
|
||||
putPooledBytes(&voicePayloadPool, voicePayload)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func (h *PacketHandler) Close() error {
|
||||
if h.closed.Swap(true) {
|
||||
return nil
|
||||
}
|
||||
close(h.stopCh)
|
||||
if h.conn != nil {
|
||||
return h.conn.Close()
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func getPooledBytes(pool *sync.Pool, size int) []byte {
|
||||
bufPtr, ok := pool.Get().(*[]byte)
|
||||
if !ok || bufPtr == nil {
|
||||
return make([]byte, size)
|
||||
}
|
||||
buf := *bufPtr
|
||||
if cap(buf) < size {
|
||||
return make([]byte, size)
|
||||
}
|
||||
|
||||
return buf[:size]
|
||||
}
|
||||
|
||||
func putPooledBytes(pool *sync.Pool, buf []byte) {
|
||||
pool.Put(&buf)
|
||||
}
|
||||
@@ -0,0 +1,585 @@
|
||||
package transport
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"io"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/honeybbq/teamspeak-go/crypto"
|
||||
)
|
||||
|
||||
// packetPipe is one leg of an in-memory datagram pair: one Write is one Read message.
|
||||
type packetPipe struct {
|
||||
recv <-chan []byte
|
||||
send chan<- []byte
|
||||
done chan struct{}
|
||||
once sync.Once
|
||||
closed atomic.Bool
|
||||
}
|
||||
|
||||
func (p *packetPipe) Read(b []byte) (int, error) {
|
||||
select {
|
||||
case data, ok := <-p.recv:
|
||||
if !ok {
|
||||
return 0, io.EOF
|
||||
}
|
||||
n := copy(b, data)
|
||||
|
||||
return n, nil
|
||||
case <-p.done:
|
||||
return 0, io.EOF
|
||||
}
|
||||
}
|
||||
|
||||
func (p *packetPipe) Write(b []byte) (int, error) {
|
||||
if p.closed.Load() {
|
||||
return 0, io.ErrClosedPipe
|
||||
}
|
||||
cp := make([]byte, len(b))
|
||||
copy(cp, b)
|
||||
select {
|
||||
case p.send <- cp:
|
||||
return len(b), nil
|
||||
case <-p.done:
|
||||
return 0, io.ErrClosedPipe
|
||||
}
|
||||
}
|
||||
|
||||
func (p *packetPipe) Close() error {
|
||||
p.once.Do(func() {
|
||||
p.closed.Store(true)
|
||||
close(p.done)
|
||||
})
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// newTestPair creates a pair of connected packetPipe endpoints.
|
||||
// clientConn is given to PacketHandler.Start(); serverConn is used by the test
|
||||
// to read what the handler sends and inject packets the handler receives.
|
||||
func newTestPair() (*packetPipe, *packetPipe) {
|
||||
toClient := make(chan []byte, 256)
|
||||
fromClient := make(chan []byte, 256)
|
||||
done := make(chan struct{})
|
||||
|
||||
clientConn := &packetPipe{recv: toClient, send: fromClient, done: done}
|
||||
serverConn := &packetPipe{recv: fromClient, send: toClient, done: done}
|
||||
|
||||
return clientConn, serverConn
|
||||
}
|
||||
|
||||
const testIdentityForHandler = "W2OSGpWxkzBPJjt8iyJFsMnqnwHCnxOlmE9gWFOFnKs=:0"
|
||||
|
||||
// These match the unexported dummy key/nonce in crypto/crypt_ops.go,
|
||||
// used for EAX encrypt/decrypt before CryptoInit completes.
|
||||
var (
|
||||
handlerTestDummyKey = []byte(`c:\windows\syste`)
|
||||
handlerTestDummyNonce = []byte(`m\firewall32.cpl`)
|
||||
)
|
||||
|
||||
func newTestHandler(t *testing.T) (*PacketHandler, *packetPipe) {
|
||||
t.Helper()
|
||||
id, err := crypto.IdentityFromString(testIdentityForHandler)
|
||||
if err != nil {
|
||||
t.Fatalf("IdentityFromString: %v", err)
|
||||
}
|
||||
tc := crypto.NewCrypt(id)
|
||||
h := NewPacketHandler(tc, slog.Default())
|
||||
clientConn, serverConn := newTestPair()
|
||||
startErr := h.Start(clientConn)
|
||||
if startErr != nil {
|
||||
t.Fatalf("Start: %v", startErr)
|
||||
}
|
||||
|
||||
return h, serverConn
|
||||
}
|
||||
|
||||
// readPacket reads the next packet from serverConn with a 2-second timeout.
|
||||
func readPacket(t *testing.T, serverConn *packetPipe) []byte {
|
||||
t.Helper()
|
||||
buf := make([]byte, 4096)
|
||||
done := make(chan []byte, 1)
|
||||
|
||||
go func() {
|
||||
n, err := serverConn.Read(buf)
|
||||
if err != nil {
|
||||
done <- nil
|
||||
|
||||
return
|
||||
}
|
||||
cp := make([]byte, n)
|
||||
copy(cp, buf[:n])
|
||||
done <- cp
|
||||
}()
|
||||
|
||||
select {
|
||||
case data := <-done:
|
||||
return data
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatalf("readPacket: timed out after 2s")
|
||||
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// buildS2CPacket constructs a raw S2C (server-to-client) packet for injection.
|
||||
// Format: [8 tag][2 ID][1 TypeFlagged][payload].
|
||||
func buildS2CPacket(tag []byte, id uint16, typeFlagged byte, payload []byte) []byte {
|
||||
raw := make([]byte, 8+3+len(payload))
|
||||
copy(raw[0:8], tag)
|
||||
binary.BigEndian.PutUint16(raw[8:10], id)
|
||||
raw[10] = typeFlagged
|
||||
copy(raw[11:], payload)
|
||||
|
||||
return raw
|
||||
}
|
||||
|
||||
// buildDummyEncryptedS2CCommand encrypts a command payload with the dummy EAX key
|
||||
// and returns the full raw S2C packet bytes.
|
||||
func buildDummyEncryptedS2CCommand(t *testing.T, pktID uint16, typeFlagged byte, payload []byte) []byte {
|
||||
t.Helper()
|
||||
s2cHeader := make([]byte, 3)
|
||||
binary.BigEndian.PutUint16(s2cHeader[0:2], pktID)
|
||||
s2cHeader[2] = typeFlagged
|
||||
|
||||
key := make([]byte, 16)
|
||||
copy(key, handlerTestDummyKey)
|
||||
eax, err := crypto.NewEAX(key)
|
||||
if err != nil {
|
||||
t.Fatalf("NewEAX: %v", err)
|
||||
}
|
||||
ciphertext, mac, err := eax.Encrypt(handlerTestDummyNonce, s2cHeader, payload)
|
||||
if err != nil {
|
||||
t.Fatalf("Encrypt: %v", err)
|
||||
}
|
||||
|
||||
raw := make([]byte, 8+3+len(ciphertext))
|
||||
copy(raw[0:8], mac)
|
||||
copy(raw[8:11], s2cHeader)
|
||||
copy(raw[11:], ciphertext)
|
||||
|
||||
return raw
|
||||
}
|
||||
|
||||
func TestHandlerStart_SendsInit1Packet(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
|
||||
pkt := readPacket(t, serverConn)
|
||||
if len(pkt) < 13+21 {
|
||||
t.Fatalf("expected at least %d bytes, got %d", 13+21, len(pkt))
|
||||
}
|
||||
|
||||
// Verify MAC is "TS3INIT1"
|
||||
if string(pkt[0:8]) != "TS3INIT1" {
|
||||
t.Errorf("expected 'TS3INIT1' MAC, got %q", pkt[0:8])
|
||||
}
|
||||
// C2S header[4] lower nibble is the packet type.
|
||||
typeByte := pkt[12] & 0x0F
|
||||
if typeByte != byte(PacketTypeInit1) {
|
||||
t.Errorf("expected PacketTypeInit1 (%d), got %d", PacketTypeInit1, typeByte)
|
||||
}
|
||||
// Payload: [4 version][1 type=0x00][...] = 21 bytes
|
||||
if len(pkt[13:]) != 21 {
|
||||
t.Errorf("expected 21-byte Init1 payload, got %d", len(pkt[13:]))
|
||||
}
|
||||
}
|
||||
|
||||
// TestHandlerReceive_Init1Packet_CallsOnPacket verifies that a server-sent Init1
|
||||
// packet is delivered to OnPacket. Responding with subsequent Init1 steps is the
|
||||
// Client's responsibility, not the PacketHandler's.
|
||||
func TestHandlerReceive_Init1Packet_CallsOnPacket(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
received := make(chan *Packet, 1)
|
||||
h.OnPacket = func(p *Packet) {
|
||||
received <- p
|
||||
}
|
||||
|
||||
step0Data := make([]byte, 21)
|
||||
step0Data[0] = 0x00
|
||||
binary.LittleEndian.PutUint32(step0Data[9:13], 0xCAFEBABE)
|
||||
|
||||
raw := buildS2CPacket(make([]byte, 8), 0, byte(PacketTypeInit1), step0Data)
|
||||
_, writeErr := serverConn.Write(raw)
|
||||
if writeErr != nil {
|
||||
t.Fatalf("Write: %v", writeErr)
|
||||
}
|
||||
|
||||
select {
|
||||
case p := <-received:
|
||||
if p.Type() != PacketTypeInit1 {
|
||||
t.Errorf("expected PacketTypeInit1, got %v", p.Type())
|
||||
}
|
||||
if len(p.Data) != 21 {
|
||||
t.Errorf("expected 21-byte payload, got %d", len(p.Data))
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Error("OnPacket not called for Init1 packet")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerReceive_PingFromServer_SendsPong(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
const pingID = uint16(42)
|
||||
pingPayload := make([]byte, 2)
|
||||
binary.BigEndian.PutUint16(pingPayload, pingID)
|
||||
|
||||
// TypeFlagged = PacketTypePing(4) | PacketFlagUnencrypted(0x80)
|
||||
// FakeSignature before CryptoInit = all zeros.
|
||||
typeFlagged := byte(PacketTypePing) | byte(PacketFlagUnencrypted)
|
||||
raw := buildS2CPacket(make([]byte, 8), pingID, typeFlagged, pingPayload)
|
||||
_, writeErr2 := serverConn.Write(raw)
|
||||
if writeErr2 != nil {
|
||||
t.Fatalf("Write ping: %v", writeErr2)
|
||||
}
|
||||
|
||||
resp := readPacket(t, serverConn)
|
||||
if len(resp) < 13+2 {
|
||||
t.Fatalf("expected pong, got %d bytes", len(resp))
|
||||
}
|
||||
if resp[12]&0x0F != byte(PacketTypePong) {
|
||||
t.Errorf("expected PacketTypePong (%d), got %d", PacketTypePong, resp[12]&0x0F)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerClose_CallsOnClosed(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
|
||||
closed := make(chan error, 1)
|
||||
h.OnClosed = func(err error) {
|
||||
closed <- err
|
||||
}
|
||||
|
||||
_ = readPacket(t, serverConn)
|
||||
_ = h.Close()
|
||||
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Error("OnClosed not called within timeout")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerReceive_CommandPacket_CallsOnPacket(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
typeFlagged := byte(PacketTypeCommand) | byte(PacketFlagNewProtocol)
|
||||
raw := buildDummyEncryptedS2CCommand(t, 0, typeFlagged, []byte("hello"))
|
||||
|
||||
received := make(chan *Packet, 1)
|
||||
h.OnPacket = func(p *Packet) {
|
||||
received <- p
|
||||
}
|
||||
|
||||
_, writeErr3 := serverConn.Write(raw)
|
||||
if writeErr3 != nil {
|
||||
t.Fatalf("Write command: %v", writeErr3)
|
||||
}
|
||||
|
||||
select {
|
||||
case p := <-received:
|
||||
if string(p.Data) != "hello" {
|
||||
t.Errorf("expected 'hello', got %q", p.Data)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Error("OnPacket not called within timeout")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHandlerReceive_FragmentedCommandPacket_Reassembles verifies that two
|
||||
// fragmented command packets are correctly reassembled before OnPacket is called.
|
||||
//
|
||||
// TeamSpeak fragmentation:
|
||||
// - First fragment: PacketFlagFragmented SET
|
||||
// - Middle fragments: PacketFlagFragmented NOT set
|
||||
// - Last fragment: PacketFlagFragmented SET ← both start and end have the flag
|
||||
func TestHandlerReceive_FragmentedCommandPacket_Reassembles(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
// Fragment 1 (start): ID=0, Command | NewProtocol | Fragmented
|
||||
f1Type := byte(PacketTypeCommand) | byte(PacketFlagNewProtocol) | byte(PacketFlagFragmented)
|
||||
// Fragment 2 (end): ID=1, Command | NewProtocol | Fragmented
|
||||
// Both first and last fragments have Fragmented set per TeamSpeak fragmentation rules.
|
||||
f2Type := byte(PacketTypeCommand) | byte(PacketFlagNewProtocol) | byte(PacketFlagFragmented)
|
||||
|
||||
received := make(chan *Packet, 1)
|
||||
h.OnPacket = func(p *Packet) {
|
||||
received <- p
|
||||
}
|
||||
|
||||
_, err1 := serverConn.Write(buildDummyEncryptedS2CCommand(t, 0, f1Type, []byte("hello")))
|
||||
if err1 != nil {
|
||||
t.Fatalf("Write fragment 1: %v", err1)
|
||||
}
|
||||
_, err2 := serverConn.Write(buildDummyEncryptedS2CCommand(t, 1, f2Type, []byte(" world")))
|
||||
if err2 != nil {
|
||||
t.Fatalf("Write fragment 2: %v", err2)
|
||||
}
|
||||
|
||||
select {
|
||||
case p := <-received:
|
||||
if string(p.Data) != "hello world" {
|
||||
t.Errorf("expected 'hello world', got %q", p.Data)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Error("OnPacket not called for reassembled packet")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHandlerSendPacket_LargeCommand_SplitsIntoFragments verifies that a
|
||||
// Command payload >487 bytes is fragmented.
|
||||
//
|
||||
// first != last → set PacketFlagFragmented (only on first and last, not middle).
|
||||
func TestHandlerSendPacket_LargeCommand_SplitsIntoFragments(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
// 975 bytes → 3 fragments: 487 + 487 + 1
|
||||
largeData := make([]byte, 975)
|
||||
sendErr := h.SendPacket(byte(PacketTypeCommand), largeData, 0)
|
||||
if sendErr != nil {
|
||||
t.Fatalf("SendPacket error: %v", sendErr)
|
||||
}
|
||||
|
||||
// Collect 3 fragments.
|
||||
pkts := make([][]byte, 0, 3)
|
||||
for range 3 {
|
||||
pkts = append(pkts, readPacket(t, serverConn))
|
||||
}
|
||||
|
||||
// C2S raw layout: [8 tag][2 pktID][2 clientID][1 TypeFlagged][ciphertext]
|
||||
// TypeFlagged byte is at index 12.
|
||||
fragFlag := byte(PacketFlagFragmented)
|
||||
// Fragment 0 (first=true, last=false): Fragmented set
|
||||
if pkts[0][12]&fragFlag == 0 {
|
||||
t.Errorf("fragment 0 should have Fragmented flag, TypeFlagged=0x%02x", pkts[0][12])
|
||||
}
|
||||
// Fragment 1 (first=false, last=false): Fragmented NOT set
|
||||
if pkts[1][12]&fragFlag != 0 {
|
||||
t.Errorf("fragment 1 (middle) should NOT have Fragmented flag, TypeFlagged=0x%02x", pkts[1][12])
|
||||
}
|
||||
// Fragment 2 (first=false, last=true): Fragmented set
|
||||
if pkts[2][12]&fragFlag == 0 {
|
||||
t.Errorf("fragment 2 (last) should have Fragmented flag, TypeFlagged=0x%02x", pkts[2][12])
|
||||
}
|
||||
}
|
||||
|
||||
// TestHandlerSendPacket_ExactBoundary_NoSplit verifies that a 487-byte payload
|
||||
// (exactly the max) is sent as a single packet.
|
||||
func TestHandlerSendPacket_ExactBoundary_NoSplit(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
sendErr := h.SendPacket(byte(PacketTypeCommand), make([]byte, 487), 0)
|
||||
if sendErr != nil {
|
||||
t.Fatalf("SendPacket: %v", sendErr)
|
||||
}
|
||||
|
||||
pkt := readPacket(t, serverConn)
|
||||
if pkt[12]&byte(PacketFlagFragmented) != 0 {
|
||||
t.Error("exact-boundary packet should not have Fragmented flag")
|
||||
}
|
||||
|
||||
// Ensure no second fragment arrives.
|
||||
select {
|
||||
case extra := <-func() chan []byte {
|
||||
ch := make(chan []byte, 1)
|
||||
go func() {
|
||||
buf := make([]byte, 4096)
|
||||
n, readErr := serverConn.Read(buf)
|
||||
if readErr == nil {
|
||||
cp := make([]byte, n)
|
||||
copy(cp, buf[:n])
|
||||
ch <- cp
|
||||
}
|
||||
}()
|
||||
|
||||
return ch
|
||||
}():
|
||||
t.Errorf("unexpected second packet: %d bytes", len(extra))
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerSendVoicePacket_Format(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
voiceData := []byte{0x01, 0x02, 0x03, 0x04}
|
||||
const codec = byte(4) // Opus Voice
|
||||
sendErr := h.SendVoicePacket(voiceData, codec)
|
||||
if sendErr != nil {
|
||||
t.Fatalf("SendVoicePacket: %v", sendErr)
|
||||
}
|
||||
|
||||
pkt := readPacket(t, serverConn)
|
||||
if len(pkt) < 13+3+len(voiceData) {
|
||||
t.Fatalf("packet too short: %d bytes", len(pkt))
|
||||
}
|
||||
|
||||
// Tag bytes 0-7 should be FakeSignature (all zeros for unused crypto state).
|
||||
fakeSig := h.TsCrypt.FakeSignature
|
||||
for i, b := range fakeSig {
|
||||
if pkt[i] != b {
|
||||
t.Errorf("tag[%d]: expected 0x%02x (FakeSignature), got 0x%02x", i, b, pkt[i])
|
||||
}
|
||||
}
|
||||
|
||||
// TypeFlagged byte 12: type = PacketTypeVoice (0), flags = Unencrypted (0x80)
|
||||
if pkt[12]&0x0F != byte(PacketTypeVoice) {
|
||||
t.Errorf("expected PacketTypeVoice (0), got %d", pkt[12]&0x0F)
|
||||
}
|
||||
if pkt[12]&byte(PacketFlagUnencrypted) == 0 {
|
||||
t.Error("voice packet should have Unencrypted flag")
|
||||
}
|
||||
|
||||
// Payload: [2 seqID][1 codec][data]
|
||||
if pkt[13+2] != codec {
|
||||
t.Errorf("expected codec=0x%02x, got 0x%02x", codec, pkt[13+2])
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerSendVoicePacket_SequenceIncreases(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
for i := range 3 {
|
||||
voiceSendErr := h.SendVoicePacket([]byte{byte(i)}, 4)
|
||||
if voiceSendErr != nil {
|
||||
t.Fatalf("SendVoicePacket[%d]: %v", i, voiceSendErr)
|
||||
}
|
||||
}
|
||||
|
||||
pkt0 := readPacket(t, serverConn)
|
||||
pkt1 := readPacket(t, serverConn)
|
||||
|
||||
seq0 := binary.BigEndian.Uint16(pkt0[13:15])
|
||||
seq1 := binary.BigEndian.Uint16(pkt1[13:15])
|
||||
if seq1 != seq0+1 {
|
||||
t.Errorf("expected seq1 = seq0+1 = %d, got %d", seq0+1, seq1)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerReceivedFinalInitAck_ClearsInitPacketCheck(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
h.mu.Lock()
|
||||
hasCheck := h.initPacketCheck != nil
|
||||
h.mu.Unlock()
|
||||
|
||||
if !hasCheck {
|
||||
t.Error("expected initPacketCheck to be set after Start()")
|
||||
}
|
||||
|
||||
h.ReceivedFinalInitAck()
|
||||
|
||||
h.mu.Lock()
|
||||
hasCheck = h.initPacketCheck != nil
|
||||
h.mu.Unlock()
|
||||
|
||||
if hasCheck {
|
||||
t.Error("expected initPacketCheck to be nil after ReceivedFinalInitAck()")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerGetWinForType(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
if w := h.getWinForType(PacketTypeCommand); w == nil {
|
||||
t.Error("expected non-nil window for PacketTypeCommand")
|
||||
}
|
||||
if w := h.getWinForType(PacketTypeCommandLow); w == nil {
|
||||
t.Error("expected non-nil window for PacketTypeCommandLow")
|
||||
}
|
||||
if w := h.getWinForType(PacketTypePing); w != nil {
|
||||
t.Errorf("expected nil window for PacketTypePing, got %v", w)
|
||||
}
|
||||
if w := h.getWinForType(PacketTypeVoice); w != nil {
|
||||
t.Errorf("expected nil window for PacketTypeVoice, got %v", w)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerCheckResends_InitPacket_ReSent(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
defer func() { _ = h.Close() }()
|
||||
// Drain the initial Init1 packet.
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
// Fast-forward initPacketCheck.lastSend so checkResends triggers a resend.
|
||||
h.mu.Lock()
|
||||
if h.initPacketCheck != nil {
|
||||
h.initPacketCheck.lastSend = time.Now().Add(-time.Second)
|
||||
}
|
||||
h.mu.Unlock()
|
||||
|
||||
h.checkResends()
|
||||
|
||||
// A resent Init1 packet should now appear on serverConn.
|
||||
resent := readPacket(t, serverConn)
|
||||
if len(resent) < 13 {
|
||||
t.Fatalf("expected resent Init1, got %d bytes", len(resent))
|
||||
}
|
||||
if resent[12]&0x0F != byte(PacketTypeInit1) {
|
||||
t.Errorf("expected PacketTypeInit1, got type %d", resent[12]&0x0F)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerCheckResends_IdleTimeout_ClosesHandler(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
closed := make(chan error, 1)
|
||||
h.OnClosed = func(err error) { closed <- err }
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
// Simulate idle timeout: set lastMessageReceived to a distant past.
|
||||
h.mu.Lock()
|
||||
h.lastMessageReceived = time.Now().Add(-(PacketTimeout + time.Second))
|
||||
h.mu.Unlock()
|
||||
|
||||
h.checkResends()
|
||||
|
||||
select {
|
||||
case <-closed:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Error("expected handler to close on idle timeout")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlerClose_IdempotentNoError(t *testing.T) {
|
||||
h, serverConn := newTestHandler(t)
|
||||
_ = readPacket(t, serverConn)
|
||||
|
||||
err := h.Close()
|
||||
if err != nil {
|
||||
t.Errorf("first Close() returned error: %v", err)
|
||||
}
|
||||
err = h.Close()
|
||||
if err != nil {
|
||||
t.Errorf("second Close() returned error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
package transport
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"time"
|
||||
)
|
||||
|
||||
type PacketType byte
|
||||
|
||||
const (
|
||||
PacketTypeVoice PacketType = 0
|
||||
PacketTypeVoiceWhisper PacketType = 1
|
||||
PacketTypeCommand PacketType = 2
|
||||
PacketTypeCommandLow PacketType = 3
|
||||
PacketTypePing PacketType = 4
|
||||
PacketTypePong PacketType = 5
|
||||
PacketTypeAck PacketType = 6
|
||||
PacketTypeAckLow PacketType = 7
|
||||
PacketTypeInit1 PacketType = 8
|
||||
)
|
||||
|
||||
type PacketFlags byte
|
||||
|
||||
const (
|
||||
PacketFlagFragmented PacketFlags = 0x10
|
||||
PacketFlagNewProtocol PacketFlags = 0x20
|
||||
PacketFlagCompressed PacketFlags = 0x40
|
||||
PacketFlagUnencrypted PacketFlags = 0x80
|
||||
)
|
||||
|
||||
type Packet struct {
|
||||
ReceivedAt time.Time
|
||||
Data []byte
|
||||
GenerationID uint32
|
||||
ID uint16
|
||||
ClientID uint16
|
||||
TypeFlagged byte
|
||||
}
|
||||
|
||||
func (p *Packet) Type() PacketType {
|
||||
return PacketType(p.TypeFlagged & 0x0F)
|
||||
}
|
||||
|
||||
func (p *Packet) Flags() PacketFlags {
|
||||
return PacketFlags(p.TypeFlagged & 0xF0)
|
||||
}
|
||||
|
||||
func (p *Packet) IsUnencrypted() bool {
|
||||
return (p.Flags() & PacketFlagUnencrypted) != 0
|
||||
}
|
||||
|
||||
func (p *Packet) BuildC2SHeader() []byte {
|
||||
header := make([]byte, 5)
|
||||
binary.BigEndian.PutUint16(header[0:2], p.ID)
|
||||
binary.BigEndian.PutUint16(header[2:4], p.ClientID)
|
||||
header[4] = p.TypeFlagged
|
||||
|
||||
return header
|
||||
}
|
||||
|
||||
func (p *Packet) ParseS2CHeader(raw []byte) {
|
||||
p.ID = binary.BigEndian.Uint16(raw[0:2])
|
||||
p.TypeFlagged = raw[2]
|
||||
}
|
||||
|
||||
func (p *Packet) ParseC2SHeader(raw []byte) {
|
||||
p.ID = binary.BigEndian.Uint16(raw[0:2])
|
||||
p.ClientID = binary.BigEndian.Uint16(raw[2:4])
|
||||
p.TypeFlagged = raw[4]
|
||||
}
|
||||
@@ -0,0 +1,175 @@
|
||||
package transport_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/honeybbq/teamspeak-go/transport"
|
||||
)
|
||||
|
||||
func TestPacketTypeExtraction(t *testing.T) {
|
||||
tests := []struct {
|
||||
typeFlagged byte
|
||||
wantType transport.PacketType
|
||||
}{
|
||||
{0x00, transport.PacketTypeVoice},
|
||||
{0x01, transport.PacketTypeVoiceWhisper},
|
||||
{0x02, transport.PacketTypeCommand},
|
||||
{0x03, transport.PacketTypeCommandLow},
|
||||
{0x04, transport.PacketTypePing},
|
||||
{0x05, transport.PacketTypePong},
|
||||
{0x06, transport.PacketTypeAck},
|
||||
{0x07, transport.PacketTypeAckLow},
|
||||
{0x08, transport.PacketTypeInit1},
|
||||
{0x82, transport.PacketTypeCommand}, // Unencrypted | Command
|
||||
{0xE2, transport.PacketTypeCommand}, // all flags | Command
|
||||
{0x88, transport.PacketTypeInit1}, // Unencrypted | Init1
|
||||
}
|
||||
for _, tt := range tests {
|
||||
p := &transport.Packet{TypeFlagged: tt.typeFlagged}
|
||||
if p.Type() != tt.wantType {
|
||||
t.Errorf("TypeFlagged=0x%02X: Type()=%d, want %d", tt.typeFlagged, p.Type(), tt.wantType)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketFlagsExtraction(t *testing.T) {
|
||||
tests := []struct {
|
||||
typeFlagged byte
|
||||
wantFlags transport.PacketFlags
|
||||
}{
|
||||
{0x10, transport.PacketFlagFragmented},
|
||||
{0x20, transport.PacketFlagNewProtocol},
|
||||
{0x40, transport.PacketFlagCompressed},
|
||||
{0x80, transport.PacketFlagUnencrypted},
|
||||
{
|
||||
0xF0,
|
||||
transport.PacketFlagFragmented | transport.PacketFlagNewProtocol |
|
||||
transport.PacketFlagCompressed | transport.PacketFlagUnencrypted,
|
||||
},
|
||||
{0x02, 0},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
p := &transport.Packet{TypeFlagged: tt.typeFlagged}
|
||||
if p.Flags() != tt.wantFlags {
|
||||
t.Errorf("TypeFlagged=0x%02X: Flags()=0x%02X, want 0x%02X", tt.typeFlagged, p.Flags(), tt.wantFlags)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketIsUnencrypted(t *testing.T) {
|
||||
tests := []struct {
|
||||
typeFlagged byte
|
||||
want bool
|
||||
}{
|
||||
{byte(transport.PacketFlagUnencrypted) | byte(transport.PacketTypeCommand), true},
|
||||
{byte(transport.PacketTypeCommand), false},
|
||||
{byte(transport.PacketFlagCompressed) | byte(transport.PacketTypeCommand), false},
|
||||
{0xFF, true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
p := &transport.Packet{TypeFlagged: tt.typeFlagged}
|
||||
if p.IsUnencrypted() != tt.want {
|
||||
t.Errorf("TypeFlagged=0x%02X: IsUnencrypted()=%v, want %v", tt.typeFlagged, p.IsUnencrypted(), tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildParseC2SHeaderRoundtrip(t *testing.T) {
|
||||
tests := []struct {
|
||||
id uint16
|
||||
clientID uint16
|
||||
typeFlagged byte
|
||||
}{
|
||||
{0x0001, 0x0001, byte(transport.PacketTypeCommand)},
|
||||
{0xFFFF, 0xFFFF, byte(transport.PacketTypeInit1) | byte(transport.PacketFlagUnencrypted)},
|
||||
{0x1234, 0x5678, byte(transport.PacketTypeVoice) | byte(transport.PacketFlagUnencrypted)},
|
||||
{0x0000, 0x0000, 0x00},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
p := &transport.Packet{ID: tt.id, ClientID: tt.clientID, TypeFlagged: tt.typeFlagged}
|
||||
header := p.BuildC2SHeader()
|
||||
if len(header) != 5 {
|
||||
t.Fatalf("C2S header len=%d, want 5", len(header))
|
||||
}
|
||||
p2 := &transport.Packet{}
|
||||
p2.ParseC2SHeader(header)
|
||||
if p2.ID != p.ID {
|
||||
t.Errorf("ID: got %d, want %d", p2.ID, p.ID)
|
||||
}
|
||||
if p2.ClientID != p.ClientID {
|
||||
t.Errorf("ClientID: got %d, want %d", p2.ClientID, p.ClientID)
|
||||
}
|
||||
if p2.TypeFlagged != p.TypeFlagged {
|
||||
t.Errorf("TypeFlagged: got 0x%02X, want 0x%02X", p2.TypeFlagged, p.TypeFlagged)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseS2CHeader(t *testing.T) {
|
||||
tests := []struct {
|
||||
raw []byte
|
||||
wantID uint16
|
||||
wantTypeFl byte
|
||||
}{
|
||||
{[]byte{0x00, 0x01, byte(transport.PacketTypeCommand)}, 1, byte(transport.PacketTypeCommand)},
|
||||
{[]byte{0xFF, 0xFF, byte(transport.PacketTypeInit1)}, 0xFFFF, byte(transport.PacketTypeInit1)},
|
||||
{[]byte{0x12, 0x34, 0x82}, 0x1234, 0x82},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
p := &transport.Packet{}
|
||||
p.ParseS2CHeader(tt.raw)
|
||||
if p.ID != tt.wantID {
|
||||
t.Errorf("S2C ID: got %d, want %d", p.ID, tt.wantID)
|
||||
}
|
||||
if p.TypeFlagged != tt.wantTypeFl {
|
||||
t.Errorf("S2C TypeFlagged: got 0x%02X, want 0x%02X", p.TypeFlagged, tt.wantTypeFl)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketTypeConstants(t *testing.T) {
|
||||
if transport.PacketTypeVoice != 0 {
|
||||
t.Error("PacketTypeVoice should be 0")
|
||||
}
|
||||
if transport.PacketTypeVoiceWhisper != 1 {
|
||||
t.Error("PacketTypeVoiceWhisper should be 1")
|
||||
}
|
||||
if transport.PacketTypeCommand != 2 {
|
||||
t.Error("PacketTypeCommand should be 2")
|
||||
}
|
||||
if transport.PacketTypeCommandLow != 3 {
|
||||
t.Error("PacketTypeCommandLow should be 3")
|
||||
}
|
||||
if transport.PacketTypePing != 4 {
|
||||
t.Error("PacketTypePing should be 4")
|
||||
}
|
||||
if transport.PacketTypePong != 5 {
|
||||
t.Error("PacketTypePong should be 5")
|
||||
}
|
||||
if transport.PacketTypeAck != 6 {
|
||||
t.Error("PacketTypeAck should be 6")
|
||||
}
|
||||
if transport.PacketTypeAckLow != 7 {
|
||||
t.Error("PacketTypeAckLow should be 7")
|
||||
}
|
||||
if transport.PacketTypeInit1 != 8 {
|
||||
t.Error("PacketTypeInit1 should be 8")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketReceivedAt(t *testing.T) {
|
||||
now := time.Now()
|
||||
p := &transport.Packet{ReceivedAt: now}
|
||||
if !p.ReceivedAt.Equal(now) {
|
||||
t.Error("ReceivedAt mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPacketDataField(t *testing.T) {
|
||||
data := []byte{0x01, 0x02, 0x03}
|
||||
p := &transport.Packet{Data: data}
|
||||
if len(p.Data) != 3 {
|
||||
t.Errorf("Data len=%d, want 3", len(p.Data))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
package transport
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
)
|
||||
|
||||
var (
|
||||
errQlzDataTooShort = errors.New("data too short")
|
||||
errQlzUnsupportedLevel = errors.New("only QuickLZ level 1 is supported")
|
||||
errQlzDataTooShortForHeader = errors.New("data too short for header")
|
||||
)
|
||||
|
||||
// TableSize is the QuickLZ level-1 hash table size.
|
||||
const TableSize = 4096
|
||||
|
||||
type Qlz struct {
|
||||
hashtable [TableSize]int
|
||||
}
|
||||
|
||||
type qlzState struct {
|
||||
control uint32
|
||||
sourcePos int
|
||||
destPos int
|
||||
nextHashed int
|
||||
}
|
||||
|
||||
func NewQlz() *Qlz {
|
||||
return &Qlz{}
|
||||
}
|
||||
|
||||
func getDecompressedSize(data []byte) int {
|
||||
if (data[0] & 0x02) != 0 {
|
||||
return int(binary.LittleEndian.Uint32(data[5:9]))
|
||||
}
|
||||
|
||||
return int(data[2])
|
||||
}
|
||||
|
||||
func (q *Qlz) Decompress(data []byte) ([]byte, error) {
|
||||
headerLen, decompressedSize, flags, err := parseQlzHeader(data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
dest := make([]byte, decompressedSize)
|
||||
|
||||
if (flags & 0x01) == 0 {
|
||||
copy(dest, data[headerLen:headerLen+decompressedSize])
|
||||
|
||||
return dest, nil
|
||||
}
|
||||
|
||||
for i := range q.hashtable {
|
||||
q.hashtable[i] = 0
|
||||
}
|
||||
|
||||
state := qlzState{
|
||||
control: 1,
|
||||
sourcePos: headerLen,
|
||||
}
|
||||
|
||||
for q.ensureControl(data, &state) {
|
||||
if (state.control & 1) != 0 {
|
||||
if !q.processReference(data, dest, &state) {
|
||||
break
|
||||
}
|
||||
} else {
|
||||
if q.processLiteral(data, dest, decompressedSize, &state) {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return dest, nil
|
||||
}
|
||||
|
||||
func parseQlzHeader(data []byte) (int, int, byte, error) {
|
||||
if len(data) < 3 {
|
||||
return 0, 0, 0, errQlzDataTooShort
|
||||
}
|
||||
flags := data[0]
|
||||
level := (flags >> 2) & 0x03
|
||||
if level != 1 {
|
||||
return 0, 0, 0, errQlzUnsupportedLevel
|
||||
}
|
||||
headerLen := 3
|
||||
if (flags & 0x02) != 0 {
|
||||
headerLen = 9
|
||||
}
|
||||
if len(data) < headerLen {
|
||||
return 0, 0, 0, errQlzDataTooShortForHeader
|
||||
}
|
||||
|
||||
return headerLen, getDecompressedSize(data), flags, nil
|
||||
}
|
||||
|
||||
func (q *Qlz) ensureControl(data []byte, st *qlzState) bool {
|
||||
if st.control != 1 {
|
||||
return true
|
||||
}
|
||||
if st.sourcePos+4 > len(data) {
|
||||
return false
|
||||
}
|
||||
st.control = binary.LittleEndian.Uint32(data[st.sourcePos : st.sourcePos+4])
|
||||
st.sourcePos += 4
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func (q *Qlz) processReference(data, dest []byte, st *qlzState) bool {
|
||||
st.control >>= 1
|
||||
if st.sourcePos+2 > len(data) {
|
||||
return false
|
||||
}
|
||||
b1 := data[st.sourcePos]
|
||||
b2 := data[st.sourcePos+1]
|
||||
st.sourcePos += 2
|
||||
|
||||
hash := int(b1>>4) | (int(b2) << 4)
|
||||
matchlen := int(b1 & 0x0F)
|
||||
if matchlen != 0 {
|
||||
matchlen += 2
|
||||
} else {
|
||||
if st.sourcePos >= len(data) {
|
||||
return false
|
||||
}
|
||||
matchlen = int(data[st.sourcePos])
|
||||
st.sourcePos++
|
||||
}
|
||||
|
||||
offset := q.hashtable[hash]
|
||||
for i := range matchlen {
|
||||
if st.destPos < len(dest) && offset+i < st.destPos {
|
||||
dest[st.destPos] = dest[offset+i]
|
||||
st.destPos++
|
||||
}
|
||||
}
|
||||
|
||||
end := st.destPos + 1 - matchlen
|
||||
q.updateHashtable(dest, &st.nextHashed, end)
|
||||
st.nextHashed = st.destPos
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func (q *Qlz) processLiteral(data, dest []byte, decompressedSize int, st *qlzState) bool {
|
||||
if st.destPos >= max(decompressedSize, 10)-10 {
|
||||
for st.destPos < decompressedSize {
|
||||
if st.control == 1 {
|
||||
st.sourcePos += 4
|
||||
if st.sourcePos > len(data) {
|
||||
break
|
||||
}
|
||||
st.control = binary.LittleEndian.Uint32(data[st.sourcePos-4 : st.sourcePos])
|
||||
}
|
||||
if st.sourcePos >= len(data) {
|
||||
break
|
||||
}
|
||||
dest[st.destPos] = data[st.sourcePos]
|
||||
st.destPos++
|
||||
st.sourcePos++
|
||||
st.control >>= 1
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
if st.sourcePos >= len(data) || st.destPos >= len(dest) {
|
||||
return true
|
||||
}
|
||||
dest[st.destPos] = data[st.sourcePos]
|
||||
st.destPos++
|
||||
st.sourcePos++
|
||||
st.control >>= 1
|
||||
end := max(st.destPos-2, 0)
|
||||
q.updateHashtable(dest, &st.nextHashed, end)
|
||||
if st.nextHashed < end {
|
||||
st.nextHashed = end
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func (q *Qlz) updateHashtable(dest []byte, nextHashed *int, end int) {
|
||||
for *nextHashed < end {
|
||||
if *nextHashed+3 > len(dest) {
|
||||
break
|
||||
}
|
||||
v := uint32(dest[*nextHashed]) | (uint32(dest[*nextHashed+1]) << 8) | (uint32(dest[*nextHashed+2]) << 16)
|
||||
hash := ((v >> 12) ^ v) & 0xFFF
|
||||
q.hashtable[hash] = *nextHashed
|
||||
*nextHashed++
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,140 @@
|
||||
package transport_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
|
||||
"github.com/honeybbq/teamspeak-go/transport"
|
||||
)
|
||||
|
||||
// buildUncompressedQlz constructs a QuickLZ level-1 frame with no compression.
|
||||
// flags=0x04: level=1 (bits 3:2=01), 3-byte header (bit 1=0), uncompressed (bit 0=0).
|
||||
func buildUncompressedQlz(payload []byte) []byte {
|
||||
if len(payload) > 252 {
|
||||
panic("buildUncompressedQlz: payload too large for single-byte header")
|
||||
}
|
||||
data := make([]byte, 3+len(payload))
|
||||
data[0] = 0x04
|
||||
data[1] = byte(3 + len(payload))
|
||||
data[2] = byte(len(payload))
|
||||
copy(data[3:], payload)
|
||||
|
||||
return data
|
||||
}
|
||||
|
||||
// buildAllLiteralQlz constructs a QuickLZ level-1 compressed frame where all
|
||||
// data is encoded as literals (no back-references). Control word=0 means all
|
||||
// 32 control bits select the literal path.
|
||||
// flags=0x05: level=1, 3-byte header, compressed.
|
||||
func buildAllLiteralQlz(payload []byte) []byte {
|
||||
if len(payload) > 248 {
|
||||
panic("buildAllLiteralQlz: payload too large for single-byte header")
|
||||
}
|
||||
data := make([]byte, 0, 3+4+len(payload))
|
||||
data = append(data, 0x05)
|
||||
data = append(data, byte(7+len(payload)))
|
||||
data = append(data, byte(len(payload)))
|
||||
data = append(data, 0x00, 0x00, 0x00, 0x00) // control word: all literals
|
||||
data = append(data, payload...)
|
||||
|
||||
return data
|
||||
}
|
||||
|
||||
func TestQlzDecompressErrorTooShort(t *testing.T) {
|
||||
q := transport.NewQlz()
|
||||
_, err := q.Decompress([]byte{0x04})
|
||||
if err == nil {
|
||||
t.Error("expected error for too-short data (< 3 bytes)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestQlzDecompressErrorWrongLevel(t *testing.T) {
|
||||
q := transport.NewQlz()
|
||||
// flags=0x08: level=(0x08>>2)&0x03=2 (unsupported)
|
||||
_, err := q.Decompress([]byte{0x08, 0x00, 0x04})
|
||||
if err == nil {
|
||||
t.Error("expected error for non-level-1 data")
|
||||
}
|
||||
}
|
||||
|
||||
func TestQlzDecompressUncompressed3ByteHeader(t *testing.T) {
|
||||
payload := []byte("ABCD")
|
||||
q := transport.NewQlz()
|
||||
result, err := q.Decompress(buildUncompressedQlz(payload))
|
||||
if err != nil {
|
||||
t.Fatalf("Decompress failed: %v", err)
|
||||
}
|
||||
if !bytes.Equal(result, payload) {
|
||||
t.Errorf("result=%v, want %v", result, payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestQlzDecompressUncompressed9ByteHeader(t *testing.T) {
|
||||
// flags=0x06: level=1, 9-byte header (bit 1=1), uncompressed (bit 0=0).
|
||||
// Decompressed size is uint32 LE at bytes [5:9].
|
||||
payload := []byte("ABCD")
|
||||
data := make([]byte, 9+len(payload))
|
||||
data[0] = 0x06
|
||||
data[5] = byte(len(payload))
|
||||
copy(data[9:], payload)
|
||||
q := transport.NewQlz()
|
||||
result, err := q.Decompress(data)
|
||||
if err != nil {
|
||||
t.Fatalf("Decompress failed: %v", err)
|
||||
}
|
||||
if !bytes.Equal(result, payload) {
|
||||
t.Errorf("result=%v, want %v", result, payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestQlzDecompressCompressedAllLiterals(t *testing.T) {
|
||||
// 13 bytes: max(13,10)-10=3 normal-literal iterations,
|
||||
// then the remaining 10 go through the near-end literal path.
|
||||
payload := []byte("Hello, World!")
|
||||
q := transport.NewQlz()
|
||||
result, err := q.Decompress(buildAllLiteralQlz(payload))
|
||||
if err != nil {
|
||||
t.Fatalf("Decompress failed: %v", err)
|
||||
}
|
||||
if !bytes.Equal(result, payload) {
|
||||
t.Errorf("result=%q, want %q", result, payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestQlzDecompressCompressedShortPayload(t *testing.T) {
|
||||
// Payload shorter than 10 bytes: all iterations use the near-end path.
|
||||
payload := []byte("Hi!")
|
||||
q := transport.NewQlz()
|
||||
result, err := q.Decompress(buildAllLiteralQlz(payload))
|
||||
if err != nil {
|
||||
t.Fatalf("Decompress failed: %v", err)
|
||||
}
|
||||
if !bytes.Equal(result, payload) {
|
||||
t.Errorf("result=%q, want %q", result, payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestQlzDecompressMultipleCalls(t *testing.T) {
|
||||
// Verify hashtable is reset between calls (no state leak)
|
||||
q := transport.NewQlz()
|
||||
for _, payload := range [][]byte{[]byte("first call data"), []byte("second call data")} {
|
||||
result, err := q.Decompress(buildAllLiteralQlz(payload))
|
||||
if err != nil {
|
||||
t.Fatalf("Decompress failed: %v", err)
|
||||
}
|
||||
if !bytes.Equal(result, payload) {
|
||||
t.Errorf("result=%q, want %q", result, payload)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestQlzDecompressEmptyPayload(t *testing.T) {
|
||||
q := transport.NewQlz()
|
||||
result, err := q.Decompress(buildUncompressedQlz([]byte{}))
|
||||
if err != nil {
|
||||
t.Fatalf("Decompress failed: %v", err)
|
||||
}
|
||||
if len(result) != 0 {
|
||||
t.Errorf("expected empty result, got %v", result)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user