Ground-Zerro / Phobos Public
Code Issues Pull requests Actions Releases View on GitHub ↗
2.4 KB go
/* SPDX-License-Identifier: MIT
 *
 * Phobos
 */

package phobos

import (
	"errors"
	"net"
	"sync"

	"golang.zx2c4.com/wireguard/windows/phobos/cobf"
)

var errFrameCorrupt = errors.New("phobos: corrupt obfuscated frame")

type obfConn struct {
	net.Conn
	masking Masking
	media   MediaParams

	writeMu      sync.Mutex
	writeCipher  *cobf.StreamCipher
	encoder      s5Encoder
	writeScratch []byte
	writeFrames  []byte

	readCipher *cobf.StreamCipher
	decoder    s5Decoder
	readRaw    []byte
	readPlain  []byte
	pending    []byte
	readErr    error
}

func newObfConn(conn net.Conn, key []byte, masking Masking, media MediaParams) *obfConn {
	c := &obfConn{
		Conn:        conn,
		masking:     masking,
		media:       media,
		writeCipher: cobf.NewStreamCipher(key),
		readCipher:  cobf.NewStreamCipher(key),
		readRaw:     make([]byte, s5EncodeReadMax),
	}
	if masking != MaskingNone {
		c.writeScratch = make([]byte, s5EncodeReadMax)
		c.writeFrames = make([]byte, s5BufferSize)
		c.readPlain = make([]byte, s5BufferSize)
	} else {
		c.writeScratch = make([]byte, s5EncodeReadMax)
	}
	return c
}

func (c *obfConn) Write(p []byte) (int, error) {
	c.writeMu.Lock()
	defer c.writeMu.Unlock()

	written := 0
	for len(p) > 0 {
		chunk := min(len(p), s5EncodeReadMax)
		plain := c.writeScratch[:chunk]
		copy(plain, p[:chunk])
		c.writeCipher.Apply(plain)

		wire := plain
		if c.masking != MaskingNone {
			n := c.encoder.encode(c.masking, c.media, plain, c.writeFrames)
			if n < 0 {
				return written, errFrameCorrupt
			}
			wire = c.writeFrames[:n]
		}
		if _, err := c.Conn.Write(wire); err != nil {
			return written, err
		}
		written += chunk
		p = p[chunk:]
	}
	return written, nil
}

func (c *obfConn) Read(p []byte) (int, error) {
	for len(c.pending) == 0 {
		if c.readErr != nil {
			return 0, c.readErr
		}
		n, err := c.Conn.Read(c.readRaw)
		if n > 0 {
			decoded, decodeErr := c.decodeChunk(c.readRaw[:n])
			if decodeErr != nil {
				return 0, decodeErr
			}
			c.pending = decoded
		}
		if err != nil {
			c.readErr = err
			if len(c.pending) == 0 {
				return 0, err
			}
		}
	}
	n := copy(p, c.pending)
	c.pending = c.pending[n:]
	return n, nil
}

func (c *obfConn) decodeChunk(raw []byte) ([]byte, error) {
	if c.masking == MaskingNone {
		c.readCipher.Apply(raw)
		return raw, nil
	}
	n := c.decoder.decode(c.masking, c.media, raw, c.readPlain)
	if n < 0 {
		return nil, errFrameCorrupt
	}
	plain := c.readPlain[:n]
	c.readCipher.Apply(plain)
	return plain, nil
}