New
This commit is contained in:
@@ -31,12 +31,17 @@ var BufferPool = sync.Pool{
|
||||
}
|
||||
|
||||
func ReadRequestFrame(r io.Reader) (uint32, uint32, []byte, error) {
|
||||
return ReadRequestFrameProfile(r, 0)
|
||||
}
|
||||
|
||||
// ReadRequestFrameProfile decodes the UP magic after applying headerMask.
|
||||
func ReadRequestFrameProfile(r io.Reader, headerMask byte) (uint32, uint32, []byte, error) {
|
||||
var header [14]byte
|
||||
if _, err := io.ReadFull(r, header[:]); err != nil {
|
||||
return 0, 0, nil, err
|
||||
}
|
||||
|
||||
if header[0] != 'U' || header[1] != 'P' {
|
||||
if header[0]^headerMask != 'U' || header[1]^headerMask != 'P' {
|
||||
return 0, 0, nil, errors.New("bad request magic")
|
||||
}
|
||||
|
||||
@@ -58,11 +63,16 @@ func ReadRequestFrame(r io.Reader) (uint32, uint32, []byte, error) {
|
||||
}
|
||||
|
||||
func WriteRequestFrame(w io.Writer, requestID uint32, payload []byte) error {
|
||||
return WriteRequestFrameProfile(w, requestID, payload, 0)
|
||||
}
|
||||
|
||||
// WriteRequestFrameProfile masks the two-byte UP magic with headerMask.
|
||||
func WriteRequestFrameProfile(w io.Writer, requestID uint32, payload []byte, headerMask byte) error {
|
||||
if len(payload) > MaxHandshake {
|
||||
return errors.New("request frame payload too large")
|
||||
}
|
||||
packet := make([]byte, 14+len(payload))
|
||||
packet[0], packet[1] = 'U', 'P'
|
||||
packet[0], packet[1] = 'U'^headerMask, 'P'^headerMask
|
||||
binary.BigEndian.PutUint32(packet[2:6], requestID)
|
||||
binary.BigEndian.PutUint32(packet[6:10], 0)
|
||||
binary.BigEndian.PutUint32(packet[10:14], uint32(len(payload)))
|
||||
@@ -72,12 +82,17 @@ func WriteRequestFrame(w io.Writer, requestID uint32, payload []byte) error {
|
||||
}
|
||||
|
||||
func ReadResponseFrame(r io.Reader) (uint32, []byte, error) {
|
||||
return ReadResponseFrameProfile(r, 0)
|
||||
}
|
||||
|
||||
// ReadResponseFrameProfile decodes the OK magic after applying headerMask.
|
||||
func ReadResponseFrameProfile(r io.Reader, headerMask byte) (uint32, []byte, error) {
|
||||
var header [10]byte
|
||||
if _, err := io.ReadFull(r, header[:]); err != nil {
|
||||
return 0, nil, err
|
||||
}
|
||||
|
||||
if header[0] != 'O' || header[1] != 'K' {
|
||||
if header[0]^headerMask != 'O' || header[1]^headerMask != 'K' {
|
||||
return 0, nil, fmt.Errorf("bad response magic: %q", header[:2])
|
||||
}
|
||||
|
||||
@@ -98,11 +113,16 @@ func ReadResponseFrame(r io.Reader) (uint32, []byte, error) {
|
||||
}
|
||||
|
||||
func WriteResponseFrame(w io.Writer, requestID uint32, payload []byte) error {
|
||||
return WriteResponseFrameProfile(w, requestID, payload, writerHeaderMask(w))
|
||||
}
|
||||
|
||||
// WriteResponseFrameProfile masks the two-byte OK magic with headerMask.
|
||||
func WriteResponseFrameProfile(w io.Writer, requestID uint32, payload []byte, headerMask byte) error {
|
||||
if len(payload) > MaxHandshake {
|
||||
return errors.New("response frame payload too large")
|
||||
}
|
||||
packet := make([]byte, 10+len(payload))
|
||||
packet[0], packet[1] = 'O', 'K'
|
||||
packet[0], packet[1] = 'O'^headerMask, 'K'^headerMask
|
||||
binary.BigEndian.PutUint32(packet[2:6], requestID)
|
||||
binary.BigEndian.PutUint32(packet[6:10], uint32(len(payload)))
|
||||
copy(packet[10:], payload)
|
||||
@@ -121,6 +141,13 @@ func writeAll(w io.Writer, b []byte) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func writerHeaderMask(w io.Writer) byte {
|
||||
if profiled, ok := w.(interface{ HeaderMask() byte }); ok {
|
||||
return profiled.HeaderMask()
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func CopyXOR(dst net.Conn, src net.Conn) error {
|
||||
ptr := BufferPool.Get().(*[]byte)
|
||||
buf := *ptr
|
||||
|
||||
@@ -0,0 +1,50 @@
|
||||
package protocol
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestXORHeaderProfilesRoundTrip(t *testing.T) {
|
||||
for n := 0; n < 256; n++ {
|
||||
mask := byte(n)
|
||||
if ('U'^mask)&7 < 5 {
|
||||
continue
|
||||
}
|
||||
|
||||
var request bytes.Buffer
|
||||
if err := WriteRequestFrameProfile(&request, 7, []byte("CPROBE -"), mask); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
requestID, _, payload, err := ReadRequestFrameProfile(&request, mask)
|
||||
if err != nil || requestID != 7 || !bytes.Equal(payload, []byte("CPROBE -")) {
|
||||
t.Fatalf("mask %02x request did not round-trip: id=%d payload=%q err=%v", mask, requestID, payload, err)
|
||||
}
|
||||
|
||||
var response bytes.Buffer
|
||||
if err := WriteResponseFrameProfile(&response, 7, []byte("PROBEOK"), mask); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
responseID, payload, err := ReadResponseFrameProfile(&response, mask)
|
||||
if err != nil || responseID != 7 || !bytes.Equal(payload, []byte("PROBEOK")) {
|
||||
t.Fatalf("mask %02x response did not round-trip: id=%d payload=%q err=%v", mask, responseID, payload, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type profiledBuffer struct {
|
||||
bytes.Buffer
|
||||
mask byte
|
||||
}
|
||||
|
||||
func (b *profiledBuffer) HeaderMask() byte { return b.mask }
|
||||
|
||||
func TestServerResponseUsesConnectionProfile(t *testing.T) {
|
||||
profiled := &profiledBuffer{mask: 0x3a}
|
||||
if err := WriteResponseFrame(profiled, 9, []byte("ok")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := profiled.Bytes()[0]; got != 'O'^profiled.mask {
|
||||
t.Fatalf("first byte=%02x, want %02x", got, byte('O')^profiled.mask)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user