package wire import ( "bytes" "net" "sync" "testing" ) // newFramedPair returns two FramedConns wired over an in-memory pipe with // matching per-direction ciphers (client out=C2S/in=S2C, server the reverse). func newFramedPair(key []byte) (client, server *FramedConn) { c, s := net.Pipe() client = NewFramedConn(c, CipherFor(key, DirS2C), CipherFor(key, DirC2S)) server = NewFramedConn(s, CipherFor(key, DirC2S), CipherFor(key, DirS2C)) return client, server } func TestVarIntRoundTrip(t *testing.T) { cases := []int{0, 1, 127, 128, 255, 300, 16384, 2097151, 1 << 30} for _, v := range cases { enc := AppendVarInt(nil, v) if len(enc) != VarIntSize(v) { t.Fatalf("size mismatch for %d: got %d want %d", v, len(enc), VarIntSize(v)) } got, err := ReadVarInt(bytes.NewReader(enc)) if err != nil { t.Fatalf("read %d: %v", v, err) } if got != v { t.Fatalf("roundtrip %d -> %d", v, got) } } } // TestPSKAddress cross-validates SHA3-224 against the value the Java hub prints // for the PSK "test-psk" (locks the two implementations together). func TestPSKAddress(t *testing.T) { const want = "90188f2d84e273e4d6fb27194b4a88ad10bcc20de00c493beae6d18f" if got := PSKAddress([]byte("test-psk")); got != want { t.Fatalf("PSKAddress = %s, want %s", got, want) } } // TestFramedConnRoundTrip exercises the encrypted framing + keystream continuity // in both directions over an in-memory pipe. func TestFramedConnRoundTrip(t *testing.T) { cli, srv := newFramedPair([]byte("unit-key")) // Frames of varying sizes to exercise partial-block keystream state. payloads := [][]byte{ []byte("a"), bytes.Repeat([]byte{0xAB}, 63), bytes.Repeat([]byte{0xCD}, 64), bytes.Repeat([]byte{0xEF}, 65), bytes.Repeat([]byte("mux"), 5000), } var wg sync.WaitGroup wg.Add(1) go func() { defer wg.Done() for _, p := range payloads { if err := cli.WriteFrame(p); err != nil { t.Errorf("client write: %v", err) return } } }() for _, want := range payloads { got, err := srv.ReadFrame() if err != nil { t.Fatalf("server read: %v", err) } if !bytes.Equal(got, want) { t.Fatalf("frame mismatch: len(got)=%d len(want)=%d", len(got), len(want)) } } wg.Wait() // Reverse direction. go func() { _ = srv.WriteFrame([]byte("pong")) }() got, err := cli.ReadFrame() if err != nil { t.Fatalf("client read: %v", err) } if !bytes.Equal(got, []byte("pong")) { t.Fatalf("reverse frame mismatch: %q", got) } } func TestProxyBufferShapeIsStable(t *testing.T) { // A frame with an empty payload must still be a valid (zero-length) frame. cli, srv := newFramedPair([]byte("k")) go func() { _ = cli.WriteFrame(nil) }() got, err := srv.ReadFrame() if err != nil { t.Fatalf("read empty frame: %v", err) } if len(got) != 0 { t.Fatalf("expected empty payload, got %d bytes", len(got)) } }