首次推送

This commit is contained in:
sansen
2026-07-20 19:01:03 +08:00
parent ea01b9cf99
commit ef1bf61f9f
4484 changed files with 937163 additions and 1 deletions
@@ -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)
}
}