Files
sing-box-extended/transport/sudoku/obfs/sudoku/packed.go

420 lines
9.6 KiB
Go

package sudoku
import (
"bufio"
"io"
"net"
"sync"
)
const (
packedProtectedPrefixBytes = 14
packedIOBufferSize = 64 * 1024
packedDecodeBufferSize = 96 * 1024
)
// PackedConn encodes traffic with the packed Sudoku layout while preserving
// the same padding model as the regular connection.
type PackedConn struct {
net.Conn
table *Table
reader *bufio.Reader
// Read-side buffers.
rawBuf []byte
pendingData pendingBuffer
// Write-side state.
writeMu sync.Mutex
writeBuf []byte
bitBuf uint64
bitCount int
// Read-side bit accumulator.
readBitBuf uint64
readBits int
// Padding selection matches Conn's threshold-based model.
rng *sudokuRand
paddingThreshold uint64
padMarker byte
padPool []byte
}
func (pc *PackedConn) CloseWrite() error {
if pc == nil || pc.Conn == nil {
return nil
}
if cw, ok := pc.Conn.(interface{ CloseWrite() error }); ok {
return cw.CloseWrite()
}
return nil
}
func (pc *PackedConn) CloseRead() error {
if pc == nil || pc.Conn == nil {
return nil
}
if cr, ok := pc.Conn.(interface{ CloseRead() error }); ok {
return cr.CloseRead()
}
return nil
}
func NewPackedConn(c net.Conn, table *Table, pMin, pMax int) *PackedConn {
localRng := newSeededRand()
pc := &PackedConn{
Conn: c,
table: table,
reader: bufio.NewReaderSize(c, packedIOBufferSize),
rawBuf: make([]byte, packedDecodeBufferSize),
pendingData: newPendingBuffer(4096),
writeBuf: make([]byte, 0, 4096),
rng: localRng,
paddingThreshold: pickPaddingThreshold(localRng, pMin, pMax),
}
if table != nil && table.layout != nil {
pc.padMarker = table.layout.padMarker
for _, b := range table.PaddingPool {
if b != pc.padMarker {
pc.padPool = append(pc.padPool, b)
}
}
}
if len(pc.padPool) == 0 {
pc.padPool = append(pc.padPool, pc.padMarker)
}
return pc
}
func (pc *PackedConn) appendForcedPadding(out []byte) []byte {
return append(out, pc.getPaddingByte())
}
func (pc *PackedConn) nextProtectedPrefixGap() int {
return 1 + pc.rng.Intn(2)
}
func (pc *PackedConn) writeProtectedPrefix(out []byte, p []byte) ([]byte, int) {
if len(p) == 0 {
return out, 0
}
limit := len(p)
if limit > packedProtectedPrefixBytes {
limit = packedProtectedPrefixBytes
}
for padCount := 0; padCount < 1+pc.rng.Intn(2); padCount++ {
out = pc.appendForcedPadding(out)
}
gap := pc.nextProtectedPrefixGap()
effective := 0
for i := 0; i < limit; i++ {
pc.bitBuf = (pc.bitBuf << 8) | uint64(p[i])
pc.bitCount += 8
for pc.bitCount >= 6 {
pc.bitCount -= 6
group := byte(pc.bitBuf >> pc.bitCount)
if pc.bitCount == 0 {
pc.bitBuf = 0
} else {
pc.bitBuf &= (1 << pc.bitCount) - 1
}
out = appendPackedGroup(out, pc.table.layout, pc.rng, pc.paddingThreshold, pc.padPool, group)
}
effective++
if effective >= gap {
out = pc.appendForcedPadding(out)
effective = 0
gap = pc.nextProtectedPrefixGap()
}
}
return out, limit
}
func appendPackedGroup(out []byte, layout *byteLayout, rng *sudokuRand, paddingThreshold uint64, padPool []byte, group byte) []byte {
if paddingThreshold != 0 {
u := rng.Uint32()
if uint64(u) < paddingThreshold {
out = append(out, padPool[fastIntnFromUint32(rng.Uint32(), len(padPool))])
}
}
return append(out, layout.encodeGroup[group&0x3F])
}
func maybeAppendPackedPadding(out []byte, rng *sudokuRand, paddingThreshold uint64, padPool []byte) []byte {
if paddingThreshold != 0 {
u := rng.Uint32()
if uint64(u) < paddingThreshold {
out = append(out, padPool[fastIntnFromUint32(rng.Uint32(), len(padPool))])
}
}
return out
}
func (pc *PackedConn) Write(p []byte) (int, error) {
if len(p) == 0 {
return 0, nil
}
if pc == nil || pc.Conn == nil || pc.table == nil || pc.table.layout == nil || pc.rng == nil || len(pc.padPool) == 0 {
return 0, io.ErrClosedPipe
}
pc.writeMu.Lock()
defer pc.writeMu.Unlock()
needed := len(p)*3/2 + 32
if pc.paddingThreshold == 0 {
needed = ((len(p)+2)/3)*4 + 32
}
if cap(pc.writeBuf) < needed {
pc.writeBuf = make([]byte, 0, needed)
}
out := pc.writeBuf[:0]
layout := pc.table.layout
rng := pc.rng
paddingThreshold := pc.paddingThreshold
padPool := pc.padPool
var prefixN int
out, prefixN = pc.writeProtectedPrefix(out, p)
i := prefixN
n := len(p)
for pc.bitCount > 0 && i < n {
b := p[i]
i++
pc.bitBuf = (pc.bitBuf << 8) | uint64(b)
pc.bitCount += 8
for pc.bitCount >= 6 {
pc.bitCount -= 6
group := byte(pc.bitBuf >> pc.bitCount)
if pc.bitCount == 0 {
pc.bitBuf = 0
} else {
pc.bitBuf &= (1 << pc.bitCount) - 1
}
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, group)
}
}
for i+11 < n {
for batch := 0; batch < 4; batch++ {
b1, b2, b3 := p[i], p[i+1], p[i+2]
i += 3
g1 := (b1 >> 2) & 0x3F
g2 := ((b1 & 0x03) << 4) | ((b2 >> 4) & 0x0F)
g3 := ((b2 & 0x0F) << 2) | ((b3 >> 6) & 0x03)
g4 := b3 & 0x3F
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, g1)
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, g2)
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, g3)
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, g4)
}
}
for i+2 < n {
b1, b2, b3 := p[i], p[i+1], p[i+2]
i += 3
g1 := (b1 >> 2) & 0x3F
g2 := ((b1 & 0x03) << 4) | ((b2 >> 4) & 0x0F)
g3 := ((b2 & 0x0F) << 2) | ((b3 >> 6) & 0x03)
g4 := b3 & 0x3F
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, g1)
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, g2)
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, g3)
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, g4)
}
for ; i < n; i++ {
b := p[i]
pc.bitBuf = (pc.bitBuf << 8) | uint64(b)
pc.bitCount += 8
for pc.bitCount >= 6 {
pc.bitCount -= 6
group := byte(pc.bitBuf >> pc.bitCount)
if pc.bitCount == 0 {
pc.bitBuf = 0
} else {
pc.bitBuf &= (1 << pc.bitCount) - 1
}
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, group)
}
}
if pc.bitCount > 0 {
group := byte(pc.bitBuf << (6 - pc.bitCount))
pc.bitBuf = 0
pc.bitCount = 0
out = appendPackedGroup(out, layout, rng, paddingThreshold, padPool, group)
out = append(out, pc.padMarker)
}
out = maybeAppendPackedPadding(out, rng, paddingThreshold, padPool)
if len(out) > 0 {
pc.writeBuf = out[:0]
return len(p), writeFull(pc.Conn, out)
}
pc.writeBuf = out[:0]
return len(p), nil
}
func (pc *PackedConn) Flush() error {
if pc == nil || pc.Conn == nil || pc.table == nil || pc.table.layout == nil || pc.rng == nil || len(pc.padPool) == 0 {
return io.ErrClosedPipe
}
pc.writeMu.Lock()
defer pc.writeMu.Unlock()
out := pc.writeBuf[:0]
if pc.bitCount > 0 {
group := byte(pc.bitBuf << (6 - pc.bitCount))
pc.bitBuf = 0
pc.bitCount = 0
out = append(out, pc.table.layout.groupByte(group&0x3F))
out = append(out, pc.padMarker)
}
out = maybeAppendPackedPadding(out, pc.rng, pc.paddingThreshold, pc.padPool)
if len(out) > 0 {
pc.writeBuf = out[:0]
return writeFull(pc.Conn, out)
}
return nil
}
func writeFull(w io.Writer, b []byte) error {
for len(b) > 0 {
n, err := w.Write(b)
if err != nil {
return err
}
if n == 0 {
return io.ErrShortWrite
}
b = b[n:]
}
return nil
}
func (pc *PackedConn) Read(p []byte) (int, error) {
if len(p) == 0 {
return 0, nil
}
if pc == nil || pc.Conn == nil || pc.reader == nil || len(pc.rawBuf) == 0 || pc.table == nil || pc.table.layout == nil {
return 0, io.ErrClosedPipe
}
if n, ok := drainPending(p, &pc.pendingData); ok {
return n, nil
}
outN := 0
for {
nr, rErr := readRawLimited(pc.Conn, pc.reader, pc.rawBuf[:packedReadSize(len(p)-outN, len(pc.rawBuf))])
if nr > 0 {
rBuf := pc.readBitBuf
rBits := pc.readBits
padMarker := pc.padMarker
layout := pc.table.layout
chunk := pc.rawBuf[:nr]
for i := 0; i < len(chunk); {
if rBits == 0 && outN+3 <= len(p) && i+3 < len(chunk) &&
layout.hintTable[chunk[i]] && layout.hintTable[chunk[i+1]] &&
layout.hintTable[chunk[i+2]] && layout.hintTable[chunk[i+3]] {
g1 := layout.decodeGroup[chunk[i]]
g2 := layout.decodeGroup[chunk[i+1]]
g3 := layout.decodeGroup[chunk[i+2]]
g4 := layout.decodeGroup[chunk[i+3]]
p[outN] = (g1 << 2) | (g2 >> 4)
p[outN+1] = (g2 << 4) | (g3 >> 2)
p[outN+2] = (g3 << 6) | g4
outN += 3
i += 4
continue
}
b := chunk[i]
i++
if !layout.hintTable[b] {
if b == padMarker {
rBuf = 0
rBits = 0
}
continue
}
group, ok := layout.decodePackedGroup(b)
if !ok {
return 0, ErrInvalidSudokuMapMiss
}
rBuf = (rBuf << 6) | uint64(group)
rBits += 6
if rBits >= 8 {
rBits -= 8
val := byte(rBuf >> rBits)
outN = appendDecodedByte(p, outN, &pc.pendingData, val)
if rBits == 0 {
rBuf = 0
} else {
rBuf &= (uint64(1) << rBits) - 1
}
}
}
pc.readBitBuf = rBuf
pc.readBits = rBits
}
if rErr != nil {
if rErr == io.EOF {
pc.readBitBuf = 0
pc.readBits = 0
}
if outN > 0 {
return outN, nil
}
if n, ok := drainPending(p, &pc.pendingData); ok {
return n, nil
}
return 0, rErr
}
if outN > 0 {
return outN, nil
}
}
}
func (pc *PackedConn) getPaddingByte() byte {
return pc.padPool[pc.rng.Intn(len(pc.padPool))]
}
func packedReadSize(decodedRemaining, maxRaw int) int {
if maxRaw <= minDecodeReadSize || decodedRemaining <= 0 {
return maxRaw
}
if decodedRemaining > (maxRaw-minDecodeReadSize)/2 {
return maxRaw
}
return decodedRemaining*2 + minDecodeReadSize
}