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< 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") } }