354 lines
9.2 KiB
Go
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")
|
|
}
|
|
}
|