Files
rdplib/core/mppc_test.go
T

354 lines
9.2 KiB
Go

package core
import (
"bytes"
"math/rand"
"testing"
)
// Token grammar reference (MS-RDPBCGR §3.1.8.4 as implemented by FreeRDP's
// libfreerdp/codec/mppc.c):
//
// literal 0x00-0x7F : "0" + 7 bits (8-bit token)
// literal 0x80-0xFF : "10" + 7 bits (9-bit token)
// copy tuple : "11" + offset prefix + length prefix
// compressedABC is the MPPC-64K encoding of "abc" (three literal tokens;
// bytes < 0x80 are transmitted as 8-bit "0"+7bits tokens).
var compressedABC = []byte{0x61, 0x62, 0x63}
// compressedABCABC is "abc" followed by a copy tuple that repeats it:
//
// "11" copy tuple, "111" offset group 1, offset 000011 (=3),
// length "0" (=3). Bits: 01100001 01100010 01100011 11111000 0110…
// → 0x61 0x62 0x63 0xF8 0x60
//
// Decoded: "abcabc".
var compressedABCABC = []byte{0x61, 0x62, 0x63, 0xF8, 0x60}
func TestMppcDecompressLiterals(t *testing.T) {
d := NewMppcDecompressor()
got, err := d.Decompress(mppcType64K|mppcCompressed, compressedABC)
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(got, []byte("abc")) {
t.Errorf("got %q, want %q", got, "abc")
}
}
func TestMppcDecompressCopyTuple(t *testing.T) {
d := NewMppcDecompressor()
got, err := d.Decompress(mppcType64K|mppcCompressed, compressedABCABC)
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(got, []byte("abcabc")) {
t.Errorf("got %q, want %q", got, "abcabc")
}
}
func TestMppcDecompressFlush(t *testing.T) {
d := NewMppcDecompressor()
// First call: populate history.
_, _ = d.Decompress(mppcType64K|mppcCompressed, compressedABC)
if d.history[0] != 'a' {
t.Fatal("history not populated after first call")
}
// Second call with FLUSH: history must be zeroed, offset restarts at 0.
_, _ = d.Decompress(mppcAtFront|mppcFlushed|mppcCompressed, compressedABC)
if d.history[0] != 'a' || d.history[1] != 'b' || d.history[2] != 'c' {
t.Errorf("unexpected history after flush: %q %q %q",
d.history[0], d.history[1], d.history[2])
}
if d.offset != 3 {
t.Errorf("offset after flush: got %d, want 3", d.offset)
}
}
func TestMppcDecompressReset(t *testing.T) {
d := NewMppcDecompressor()
_, _ = d.Decompress(mppcType64K|mppcCompressed, compressedABC)
if d.offset != 3 {
t.Fatalf("offset after first call: got %d, want 3", d.offset)
}
// RESET (PACKET_AT_FRONT): cursor returns to the front, history contents
// preserved; the fresh tokens overwrite from position 0.
_, _ = d.Decompress(mppcAtFront|mppcType64K|mppcCompressed, compressedABC)
if d.offset != 3 {
t.Fatalf("offset after reset+decompress: got %d, want 3", d.offset)
}
if d.history[0] != 'a' {
t.Errorf("history[0] after reset: got %q", d.history[0])
}
}
func TestMppcDecompressUncompressed(t *testing.T) {
d := NewMppcDecompressor()
plain := []byte("hello")
got, err := d.Decompress(mppcType64K, plain) // no COMPRESSED flag
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(got, plain) {
t.Errorf("got %q, want %q", got, plain)
}
// Uncompressed segments do NOT advance the shared history (matching
// FreeRDP): only compressed tokens feed it.
if d.offset != 0 {
t.Errorf("offset: got %d, want 0 (history untouched)", d.offset)
}
}
// mppcTestEncoder is a minimal MPPC compressor used to generate reference
// streams. It mirrors the decoder's grammar and history arithmetic exactly.
type mppcTestEncoder struct {
buf bitWriter
hist []byte
mask int
big bool
}
type bitWriter struct {
out []byte
cur byte
nbit uint
}
func (w *bitWriter) writeBit(b int) {
if b != 0 {
w.cur |= 0x80 >> w.nbit
}
w.nbit++
if w.nbit == 8 {
w.out = append(w.out, w.cur)
w.cur = 0
w.nbit = 0
}
}
func (w *bitWriter) writeBits(v, n int) {
for i := n - 1; i >= 0; i-- {
w.writeBit((v >> uint(i)) & 1)
}
}
func (w *bitWriter) flush() []byte {
if w.nbit > 0 {
w.out = append(w.out, w.cur)
w.cur = 0
w.nbit = 0
}
return w.out
}
func newMppcTestEncoder(big bool) *mppcTestEncoder {
mask := 8191
if big {
mask = 65535
}
return &mppcTestEncoder{big: big, mask: mask}
}
func (e *mppcTestEncoder) emitLiteral(b byte) {
if b < 0x80 {
e.buf.writeBits(int(b), 8)
} else {
e.buf.writeBits(2, 2) // "10"
e.buf.writeBits(int(b)&0x7F, 7)
}
e.hist = append(e.hist, b)
}
func (e *mppcTestEncoder) emitCopy(off, length int) {
e.buf.writeBits(3, 2) // "11"
if e.big {
switch {
case off <= 63:
e.buf.writeBits(7, 3) // "111"
e.buf.writeBits(off, 6)
case off <= 319:
e.buf.writeBits(6, 3) // "110"
e.buf.writeBits(off-64, 8)
case off <= 2367:
e.buf.writeBits(2, 2) // "10"
e.buf.writeBits(off-320, 11)
default:
e.buf.writeBit(0)
e.buf.writeBits(off-2368, 16)
}
} else {
switch {
case off <= 63:
e.buf.writeBits(3, 2) // "11"
e.buf.writeBits(off, 6)
case off <= 319:
e.buf.writeBits(2, 2) // "10"
e.buf.writeBits(off-64, 8)
default:
e.buf.writeBit(0)
e.buf.writeBits(off-320, 13)
}
}
// Length: "0" → 3; otherwise n leading 1-bits, a 0 terminator, then n+1
// value bits with the (1<<m) high bit implied (m = n+1).
if length == 3 {
e.buf.writeBit(0)
} else {
m := 0
for (1 << uint(m+1)) <= length {
m++
}
k := m - 1
for i := 0; i < k; i++ {
e.buf.writeBit(1)
}
e.buf.writeBit(0)
e.buf.writeBits(length-(1<<uint(m)), m)
}
// Apply the copy to the encoder's own history view (distance = off).
pos := len(e.hist)
src := (pos - off) & e.mask
for i := 0; i < length; i++ {
e.hist = append(e.hist, e.hist[src])
src = (src + 1) & e.mask
}
}
func TestMppcLiteralRoundTrip(t *testing.T) {
rng := rand.New(rand.NewSource(1))
for _, big := range []bool{false, true} {
enc := newMppcTestEncoder(big)
var want []byte
for i := 0; i < 500; i++ {
b := byte(rng.Intn(256))
enc.emitLiteral(b)
want = append(want, b)
}
// 字面量语法与字典大小无关,两种编码器的字面量流都能被
// 64K 解码器还原。
got, err := NewMppcDecompressor().Decompress(mppcCompressed, enc.buf.flush())
if err != nil {
t.Fatalf("big=%v: %v", big, err)
}
if !bytes.Equal(got, want) {
t.Fatalf("big=%v: mismatch got %d bytes want %d", big, len(got), len(want))
}
}
}
func TestMppcCopyRoundTrip(t *testing.T) {
rng := rand.New(rand.NewSource(2))
// (offset, length) pairs covering every offset group and several length
// groups (3, 4-7, 8-15, 16-31, longer binary forms).
cases := [][2]int{
{5, 3}, {63, 4}, {64, 7}, {100, 3}, {319, 8},
{320, 16}, {1000, 31}, {2367, 9}, {2368, 32},
{5000, 100}, {30000, 300}, {60000, 500},
}
// 解码器固定使用 RDP5 64K 字典(协商决定),8K 语法不再可解。
big := true
enc := newMppcTestEncoder(big)
var want []byte
for _, c := range cases {
off, length := c[0], c[1]
for off > len(enc.hist) {
b := byte(rng.Intn(256))
enc.emitLiteral(b)
want = append(want, b)
}
enc.emitCopy(off, length)
pos := len(want)
src := (pos - off) & enc.mask
for i := 0; i < length; i++ {
want = append(want, want[src])
src = (src + 1) & enc.mask
}
}
got, err := NewMppcDecompressor().Decompress(mppcType64K|mppcCompressed, enc.buf.flush())
if err != nil {
t.Fatalf("%v", err)
}
if !bytes.Equal(got, want) {
t.Fatalf("copy round-trip mismatch (%d vs %d bytes)", len(got), len(want))
}
}
func TestMppcHistoryAcrossSegments(t *testing.T) {
enc := newMppcTestEncoder(true)
d := NewMppcDecompressor()
flags := byte(mppcCompressed) // 无 RESET/AT_FRONT:历史跨段连续
// Segment 1: seed literals.
var want []byte
for i := 0; i < 400; i++ {
b := byte('A' + i%26)
enc.emitLiteral(b)
want = append(want, b)
}
got, err := d.Decompress(flags, enc.buf.flush())
if err != nil || !bytes.Equal(got, want[:len(got)]) {
t.Fatalf("seg1: err=%v len=%d", err, len(got))
}
enc.buf.out = nil
// Segment 2: copy referencing segment 1 content.
enc.emitCopy(200, 50)
got2, err := d.Decompress(flags, enc.buf.flush())
if err != nil {
t.Fatalf("seg2: %v", err)
}
want2 := want[len(want)-200:]
for i := 0; i < 50; i++ {
if got2[i] != want2[i] {
t.Fatalf("seg2 mismatch at %d: got %q want %q", i, got2[i], want2[i])
}
}
}
func TestMppcFlushAndReset(t *testing.T) {
enc := newMppcTestEncoder(true)
d := NewMppcDecompressor()
flags := byte(mppcCompressed) // 无 RESET/AT_FRONT:历史跨段连续
for i := 0; i < 100; i++ {
enc.emitLiteral(byte(i % 256))
}
if _, err := d.Decompress(flags, enc.buf.flush()); err != nil {
t.Fatal(err)
}
if d.offset == 0 {
t.Fatal("history should have advanced")
}
// RESET moves the write cursor to the front without clearing content.
enc2 := newMppcTestEncoder(true)
enc2.emitLiteral('x')
enc2.emitLiteral('y')
got, err := d.Decompress(mppcCompressed|mppcAtFront|mppcType64K, enc2.buf.flush())
if err != nil || string(got) != "xy" {
t.Fatalf("reset: err=%v got=%q", err, got)
}
if d.offset != 2 {
t.Fatalf("reset should restart the cursor, got %d", d.offset)
}
// FLUSH clears the whole dictionary.
got, err = d.Decompress(mppcCompressed|mppcAtFront|mppcFlushed, enc2.buf.flush())
if err != nil || string(got) != "xy" {
t.Fatalf("flush: err=%v got=%q", err, got)
}
}
func TestMppcRejectsGarbage(t *testing.T) {
d := NewMppcDecompressor()
// A long run of 1-bits overflows the LengthOfMatch prefix code.
data := make([]byte, 16)
for i := range data {
data[i] = 0xFF
}
if _, err := d.Decompress(mppcCompressed|mppcType64K, data); err == nil {
t.Fatal("expected error for length code overflow")
}
}