fix: sync same-port transport detection
This commit is contained in:
parent
840ef6237c
commit
2289a31f46
10 changed files with 851 additions and 83 deletions
|
|
@ -89,7 +89,7 @@ func connLocal(conn net.Conn) string {
|
|||
|
||||
func intakeTransport(obfuscated bool) string {
|
||||
if obfuscated {
|
||||
return "obfuscated_tcp"
|
||||
return "tcp_auto"
|
||||
}
|
||||
return "tcp"
|
||||
}
|
||||
|
|
|
|||
|
|
@ -51,7 +51,13 @@ func (a frameBudgetTestAddr) String() string { return string(a) }
|
|||
|
||||
func newFrameBudgetTestTransport(packet []byte, c transport.Codec, budget *inboundFrameBudget) (*compatTransportConn, *frameBudgetTestConn) {
|
||||
raw := newFrameBudgetTestConn(packet)
|
||||
return &compatTransportConn{conn: raw, codec: c, budget: budget}, raw
|
||||
return &compatTransportConn{
|
||||
conn: raw,
|
||||
codec: c,
|
||||
codecKind: classifyInboundFrameCodec(c),
|
||||
budgetedCodec: unwrapInboundFrameBudgetedCodec(c),
|
||||
budget: budget,
|
||||
}, raw
|
||||
}
|
||||
|
||||
func TestInboundFrameBudgetSupportsBuiltInCodecs(t *testing.T) {
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
package mtprotoedge
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
|
|
@ -195,10 +194,9 @@ func (m *samePortMux) dispatch(ctx context.Context, conn net.Conn) {
|
|||
return
|
||||
}
|
||||
|
||||
wrapped := &prefixedNetConn{
|
||||
Conn: conn,
|
||||
reader: io.MultiReader(bytes.NewReader(header[:]), conn),
|
||||
}
|
||||
// Fixed eight-byte replay storage avoids bytes.Reader + MultiReader allocations
|
||||
// on every accepted same-port connection.
|
||||
wrapped := newReplayNetConn(conn, header[:])
|
||||
|
||||
target := m.tcp
|
||||
transport := "tcp"
|
||||
|
|
@ -331,17 +329,6 @@ func isSamePortMuxClosed(ch <-chan struct{}) bool {
|
|||
}
|
||||
}
|
||||
|
||||
// prefixedNetConn 把被窥探掉的前缀字节回放在数据流最前面,使下游(去混淆/codec 探测/
|
||||
// http.Server)看到完整原始字节流。
|
||||
type prefixedNetConn struct {
|
||||
reader io.Reader
|
||||
net.Conn
|
||||
}
|
||||
|
||||
func (p *prefixedNetConn) Read(b []byte) (int, error) {
|
||||
return p.reader.Read(b)
|
||||
}
|
||||
|
||||
// samePortMuxListener 是一个内存 listener:dispatch 把分流后的连接投递进来,下游
|
||||
// (serveMixed 的 accept 循环 / http.Server) 从这里 Accept。
|
||||
type samePortMuxListener struct {
|
||||
|
|
|
|||
|
|
@ -221,12 +221,14 @@ type Options struct {
|
|||
// codec 必须是 gotd 内置四种 codec(可包 NoHeader),或实现 InboundFrameBudgetedCodec;
|
||||
// 无法在 payload 分配前预检长度的 codec 会 fail-closed。
|
||||
Codec func() transport.Codec
|
||||
// ObfuscatedTCP 先按 MTProto TCP obfuscation 解包,再自动探测 codec。
|
||||
// Telegram Desktop 的 tcpo_only endpoint 会走这个 64 字节前缀流程。
|
||||
// ObfuscatedTCP 允许裸 TCP 使用 MTProto transport obfuscation。开启时按每条
|
||||
// 物理连接的首 1/4/8 字节自动区分明文 transport 与 64-byte obfuscated2,
|
||||
// 随后把 wire mode + codec 冻结到同一条双向连接;Telegram Desktop 的
|
||||
// tcpo_only 与未开启混淆的第三方客户端可共用同一端口。
|
||||
ObfuscatedTCP bool
|
||||
// WebSocket 在同一个 listener 上接受 MTProto over WebSocket(/apiws*)。
|
||||
// 开启后仅在连接建立时读取前 4 字节做 HTTP/TCP 分流;MTProto TCP
|
||||
// 后续仍走原 ObfuscatedTCP + codec 热路径。
|
||||
// 后续仍走原 TCP wire-mode + codec 探测路径。
|
||||
WebSocket bool
|
||||
// WebSocketAllowedOrigins 是允许浏览器发起 WebSocket upgrade 的页面 origin。
|
||||
// 空列表表示只接受无 Origin 的非浏览器客户端;"*" 表示允许所有来源(仅调试)。
|
||||
|
|
@ -644,7 +646,11 @@ func (s *Server) serveTCP(ctx context.Context, ln net.Listener) error {
|
|||
ctx, cancel := context.WithCancel(ctx)
|
||||
defer cancel()
|
||||
|
||||
s.log.Info("Serving", zap.String("addr", ln.Addr().String()), zap.Int("dc", s.dc), zap.Bool("obfuscated_tcp", s.obfuscated))
|
||||
s.log.Info("Serving",
|
||||
zap.String("addr", ln.Addr().String()),
|
||||
zap.Int("dc", s.dc),
|
||||
zap.String("tcp_transport_mode", intakeTransport(s.obfuscated)),
|
||||
)
|
||||
defer s.log.Info("Stopped")
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
|
|
@ -680,7 +686,7 @@ func (s *Server) serveMixed(ctx context.Context, ln net.Listener) error {
|
|||
s.log.Info("Serving",
|
||||
zap.String("addr", ln.Addr().String()),
|
||||
zap.Int("dc", s.dc),
|
||||
zap.Bool("obfuscated_tcp", s.obfuscated),
|
||||
zap.String("tcp_transport_mode", intakeTransport(s.obfuscated)),
|
||||
zap.Bool("websocket", true),
|
||||
zap.Strings("websocket_origins", s.websocketOrigins),
|
||||
)
|
||||
|
|
@ -705,7 +711,7 @@ func (s *Server) serveMixed(ctx context.Context, ln net.Listener) error {
|
|||
defer wg.Done()
|
||||
errCh <- mux.Serve(ctx)
|
||||
}()
|
||||
// 裸 MTProto TCP:每条连接在自己的 goroutine 里完成去混淆 + codec 探测。
|
||||
// 裸 MTProto TCP:每条连接在自己的 goroutine 里完成 wire mode + codec 探测。
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
errCh <- s.acceptLoop(ctx, mux.TCP(), s.obfuscated)
|
||||
|
|
@ -746,11 +752,11 @@ func (s *Server) serveMixed(ctx context.Context, ln net.Listener) error {
|
|||
return firstErr
|
||||
}
|
||||
|
||||
// acceptLoop 接受裸连接,并为每条连接单独起 goroutine 完成「去混淆 + codec 探测 +
|
||||
// acceptLoop 接受裸连接,并为每条连接单独起 goroutine 完成「wire mode + codec 探测 +
|
||||
// serveConn」。探测在 accept 循环之外、带握手超时进行——慢/半开/坏 init 的客户端只占用
|
||||
// 自己的 goroutine,绝不阻塞其他连接的接入;单条连接的握手失败也只关闭该连接,不会拖垮
|
||||
// 整个监听循环。obfuscated 为 true 时先走 obfuscated2 去混淆(裸 MTProto TCP);WebSocket
|
||||
// 连接传 false(gotd 升级处理器已完成去混淆)。
|
||||
// 整个监听循环。obfuscated 为 true 时自动区分 plain 与 obfuscated2(裸 MTProto TCP);
|
||||
// WebSocket 连接传 false(gotd 升级处理器已完成去混淆)。
|
||||
func (s *Server) acceptLoop(ctx context.Context, ln net.Listener, obfuscated bool) error {
|
||||
return s.acceptLoopTransport(ctx, ln, obfuscated, intakeTransport(obfuscated))
|
||||
}
|
||||
|
|
@ -811,7 +817,7 @@ func (s *Server) acceptLoopTransport(ctx context.Context, ln net.Listener, obfus
|
|||
func (s *Server) serveDetectedConn(ctx context.Context, raw net.Conn, obfuscated bool, transportName string) {
|
||||
started := time.Now()
|
||||
remote, local := connRemote(raw), connLocal(raw)
|
||||
// 握手读超时只覆盖去混淆 + codec 探测这一小段;用真实墙钟时间(SetReadDeadline 语义),
|
||||
// 握手读超时只覆盖 wire-mode + codec 探测这一小段;用真实墙钟时间(SetReadDeadline 语义),
|
||||
// 不走可能被测试注入的逻辑 clock。
|
||||
if err := raw.SetReadDeadline(time.Now().Add(s.handshakeTimeout)); err != nil {
|
||||
_ = raw.Close()
|
||||
|
|
@ -830,7 +836,10 @@ func (s *Server) serveDetectedConn(ctx context.Context, raw net.Conn, obfuscated
|
|||
}
|
||||
}()
|
||||
|
||||
conn, err := s.promoteConn(raw, obfuscated)
|
||||
conn, detectedTransport, err := s.promoteConn(raw, obfuscated)
|
||||
if detectedTransport != "" {
|
||||
transportName = detectedTransport
|
||||
}
|
||||
close(promoted)
|
||||
if err != nil {
|
||||
outcome := "error"
|
||||
|
|
@ -867,15 +876,15 @@ func (s *Server) serveDetectedConn(ctx context.Context, raw net.Conn, obfuscated
|
|||
}
|
||||
}
|
||||
|
||||
// promoteConn 复用与 listener 组合完全一致的「obfuscated2 去混淆 + codec 探测」管线,但针对
|
||||
// 单条连接,使其可在 accept 循环之外执行。obfuscated 对 WebSocket 连接必须为 false(gotd
|
||||
// 升级处理器已剥离 obfuscated2 并补回 codec tag)。
|
||||
func (s *Server) promoteConn(raw net.Conn, obfuscated bool) (transport.Conn, error) {
|
||||
var ln net.Listener = newSingleConnListener(raw)
|
||||
// promoteConn 针对单条连接执行一次 transport 提升。obfuscated=true 表示裸 TCP
|
||||
// 允许混淆并自动区分 plain/obfuscated2;WebSocket 必须传 false,因为 gotd upgrade
|
||||
// handler 已剥离 obfuscated2 并补回 codec tag。
|
||||
func (s *Server) promoteConn(raw net.Conn, obfuscated bool) (transport.Conn, string, error) {
|
||||
if obfuscated {
|
||||
ln = transport.ObfuscatedListener(ln)
|
||||
return s.promoteMixedTCP(raw)
|
||||
}
|
||||
return newCompatTransportListener(s.codec, ln, s.frameBudget).Accept()
|
||||
conn, err := newCompatTransportConn(s.codec, raw, s.frameBudget)
|
||||
return conn, "", err
|
||||
}
|
||||
|
||||
// serveConn 处理单个传输连接:读帧并按 auth_key_id 分流。
|
||||
|
|
|
|||
|
|
@ -232,13 +232,13 @@ func TestServerAcceptObfuscatedAbridgedQuickAckFrame(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestServerSamePortWebSocketAndObfuscatedTCP(t *testing.T) {
|
||||
func TestServerSamePortWebSocketAndMixedTCP(t *testing.T) {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
|
||||
frames := make(chan int, 2)
|
||||
frames := make(chan int, 3)
|
||||
srv := New(Options{Logger: zaptest.NewLogger(t), ObfuscatedTCP: true, WebSocket: true})
|
||||
srv.onFrame = func(n int) {
|
||||
select {
|
||||
|
|
@ -302,6 +302,27 @@ func TestServerSamePortWebSocketAndObfuscatedTCP(t *testing.T) {
|
|||
expectFrameLen(t, frames, tcpPayload.Len())
|
||||
_ = tcpConn.Close()
|
||||
|
||||
plainRaw, err := net.Dial("tcp", ln.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatalf("plain tcp dial: %v", err)
|
||||
}
|
||||
plainConn, err := transport.Intermediate.Handshake(plainRaw)
|
||||
if err != nil {
|
||||
_ = plainRaw.Close()
|
||||
t.Fatalf("plain tcp transport handshake: %v", err)
|
||||
}
|
||||
var plainPayload bin.Buffer
|
||||
plainPayload.PutInt32(0x33445566)
|
||||
plainPayload.PutInt32(0x77889900)
|
||||
sendCtx, sc = context.WithTimeout(context.Background(), 5*time.Second)
|
||||
if err := plainConn.Send(sendCtx, &plainPayload); err != nil {
|
||||
sc()
|
||||
t.Fatalf("plain tcp send: %v", err)
|
||||
}
|
||||
sc()
|
||||
expectFrameLen(t, frames, plainPayload.Len())
|
||||
_ = plainConn.Close()
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case err := <-serveErr:
|
||||
|
|
|
|||
|
|
@ -89,8 +89,7 @@ func newCompatTransportListener(codec func() transport.Codec, listener net.Liste
|
|||
}
|
||||
|
||||
// singleConnListener 是一个只产出一条「已接受」连接、随后阻塞到关闭的 net.Listener。
|
||||
// 它让单条裸连接可以走 listener 形态的去混淆/codec 管线(ObfuscatedListener +
|
||||
// compatTransportListener),从而把这部分阻塞读取从 accept 循环挪到每连接 goroutine。
|
||||
// 生产接入已直接提升单连接;这个适配器仅保留给仍需 listener 形态的测试/调用方。
|
||||
type singleConnListener struct {
|
||||
addr net.Addr
|
||||
ch chan net.Conn
|
||||
|
|
@ -131,38 +130,13 @@ func (l *compatTransportListener) Accept() (_ transport.Conn, rErr error) {
|
|||
}
|
||||
}()
|
||||
|
||||
var (
|
||||
connCodec transport.Codec
|
||||
reader io.Reader = conn
|
||||
)
|
||||
if l.codec != nil {
|
||||
connCodec = l.codec()
|
||||
if classifyInboundFrameCodec(connCodec) == inboundFrameCodecUnknown {
|
||||
// Unknown codecs are rejected before their header or first frame is read. Without an
|
||||
// explicit preflight contract, calling Codec.Read could allocate from an attacker-
|
||||
// controlled length before the process-wide budget can be reserved.
|
||||
return nil, errInboundFrameCodecUnsupported
|
||||
}
|
||||
if err := connCodec.ReadHeader(conn); err != nil {
|
||||
return nil, errors.Wrap(err, "read codec header")
|
||||
}
|
||||
} else {
|
||||
var err error
|
||||
connCodec, reader, err = detectCompatCodec(conn)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "detect codec")
|
||||
}
|
||||
promoted, err := newCompatTransportConn(l.codec, conn, l.budget)
|
||||
if err != nil {
|
||||
// Avoid returning a typed nil *compatTransportConn as a non-nil
|
||||
// transport.Conn interface on admission failure.
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &compatTransportConn{
|
||||
conn: wrappedCompatConn{
|
||||
reader: reader,
|
||||
Conn: conn,
|
||||
},
|
||||
codec: connCodec,
|
||||
budget: l.budget,
|
||||
transportPacketMessages: isTransportPacketMessageConn(conn),
|
||||
}, nil
|
||||
return promoted, nil
|
||||
}
|
||||
|
||||
func isTransportPacketMessageConn(conn net.Conn) bool {
|
||||
|
|
@ -187,10 +161,87 @@ func (w wrappedCompatConn) Read(p []byte) (int, error) {
|
|||
return w.reader.Read(p)
|
||||
}
|
||||
|
||||
// newCompatTransportConn promotes an already accepted stream without allocating the
|
||||
// single-connection listener/channel used by the historical listener composition.
|
||||
// Detection happens once per physical connection; the selected codec remains bound to
|
||||
// the returned transport for both reads and writes.
|
||||
func newCompatTransportConn(
|
||||
codecFactory func() transport.Codec,
|
||||
conn net.Conn,
|
||||
budget *inboundFrameBudget,
|
||||
) (*compatTransportConn, error) {
|
||||
if budget == nil {
|
||||
panic("mtprotoedge: nil inbound frame budget")
|
||||
}
|
||||
var (
|
||||
connCodec transport.Codec
|
||||
reader io.Reader = conn
|
||||
)
|
||||
if codecFactory != nil {
|
||||
connCodec = codecFactory()
|
||||
if classifyInboundFrameCodec(connCodec) == inboundFrameCodecUnknown {
|
||||
// Unknown codecs are rejected before their header or first frame is read. Without an
|
||||
// explicit preflight contract, calling Codec.Read could allocate from an attacker-
|
||||
// controlled length before the process-wide budget can be reserved.
|
||||
return nil, errInboundFrameCodecUnsupported
|
||||
}
|
||||
if err := connCodec.ReadHeader(conn); err != nil {
|
||||
return nil, errors.Wrap(err, "read codec header")
|
||||
}
|
||||
} else {
|
||||
var err error
|
||||
connCodec, reader, err = detectCompatCodec(conn)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "detect codec")
|
||||
}
|
||||
}
|
||||
|
||||
return newCompatTransportConnWithCodec(
|
||||
wrappedCompatConn{reader: reader, Conn: conn},
|
||||
connCodec,
|
||||
budget,
|
||||
isTransportPacketMessageConn(conn),
|
||||
)
|
||||
}
|
||||
|
||||
// newCompatTransportConnWithCodec binds a codec whose client-side transport header
|
||||
// was already consumed by the one-time wire detector.
|
||||
func newCompatTransportConnWithCodec(
|
||||
conn net.Conn,
|
||||
connCodec transport.Codec,
|
||||
budget *inboundFrameBudget,
|
||||
transportPacketMessages bool,
|
||||
) (*compatTransportConn, error) {
|
||||
if budget == nil {
|
||||
panic("mtprotoedge: nil inbound frame budget")
|
||||
}
|
||||
codecKind := classifyInboundFrameCodec(connCodec)
|
||||
if codecKind == inboundFrameCodecUnknown {
|
||||
return nil, errInboundFrameCodecUnsupported
|
||||
}
|
||||
var budgetedCodec InboundFrameBudgetedCodec
|
||||
if codecKind == inboundFrameCodecCustom {
|
||||
budgetedCodec = unwrapInboundFrameBudgetedCodec(connCodec)
|
||||
if budgetedCodec == nil {
|
||||
return nil, errInboundFrameCodecUnsupported
|
||||
}
|
||||
}
|
||||
return &compatTransportConn{
|
||||
conn: conn,
|
||||
codec: connCodec,
|
||||
codecKind: codecKind,
|
||||
budgetedCodec: budgetedCodec,
|
||||
budget: budget,
|
||||
transportPacketMessages: transportPacketMessages,
|
||||
}, nil
|
||||
}
|
||||
|
||||
type compatTransportConn struct {
|
||||
conn net.Conn
|
||||
codec transport.Codec
|
||||
budget *inboundFrameBudget
|
||||
conn net.Conn
|
||||
codec transport.Codec
|
||||
codecKind inboundFrameCodecKind
|
||||
budgetedCodec InboundFrameBudgetedCodec
|
||||
budget *inboundFrameBudget
|
||||
|
||||
transportPacketMessages bool
|
||||
directMessageScratch []byte
|
||||
|
|
@ -427,7 +478,10 @@ func (c *compatTransportConn) Close() error {
|
|||
}
|
||||
|
||||
func (c *compatTransportConn) readInboundFrame(b *bin.Buffer) error {
|
||||
kind := classifyInboundFrameCodec(c.codec)
|
||||
// codecKind and budgetedCodec are frozen when the physical connection is
|
||||
// promoted. The per-frame hot path never re-detects wire mode or type-switches
|
||||
// the already selected codec.
|
||||
kind := c.codecKind
|
||||
if kind == inboundFrameCodecUnknown {
|
||||
return errInboundFrameCodecUnsupported
|
||||
}
|
||||
|
|
@ -447,11 +501,10 @@ func (c *compatTransportConn) readInboundFrame(b *bin.Buffer) error {
|
|||
|
||||
var err error
|
||||
if kind == inboundFrameCodecCustom {
|
||||
custom := unwrapInboundFrameBudgetedCodec(c.codec)
|
||||
if custom == nil {
|
||||
if c.budgetedCodec == nil {
|
||||
return errInboundFrameCodecUnsupported
|
||||
}
|
||||
err = custom.ReadWithInboundFrameBudget(c.conn, b, reserve)
|
||||
err = c.budgetedCodec.ReadWithInboundFrameBudget(c.conn, b, reserve)
|
||||
} else {
|
||||
preflight := &inboundFramePreflightReader{r: c.conn, kind: kind, reserve: reserve}
|
||||
err = c.codec.Read(preflight, b)
|
||||
|
|
|
|||
234
internal/mtprotoedge/transport_detection.go
Normal file
234
internal/mtprotoedge/transport_detection.go
Normal file
|
|
@ -0,0 +1,234 @@
|
|||
package mtprotoedge
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
|
||||
"github.com/iamxvbaba/td/mtproxy/obfuscated2"
|
||||
"github.com/iamxvbaba/td/proto/codec"
|
||||
"github.com/iamxvbaba/td/transport"
|
||||
)
|
||||
|
||||
var errInvalidTCPTransportPrefix = errors.New("invalid MTProto TCP transport prefix")
|
||||
var errInvalidObfuscatedProtocol = errors.New("invalid obfuscated MTProto protocol tag")
|
||||
|
||||
type tcpWireMode uint8
|
||||
|
||||
const (
|
||||
tcpWireModeUnknown tcpWireMode = iota
|
||||
tcpWireModePlain
|
||||
tcpWireModeObfuscated
|
||||
)
|
||||
|
||||
type detectedTCPCodec uint8
|
||||
|
||||
const (
|
||||
detectedTCPCodecUnknown detectedTCPCodec = iota
|
||||
detectedTCPCodecAbridged
|
||||
detectedTCPCodecIntermediate
|
||||
detectedTCPCodecPaddedIntermediate
|
||||
detectedTCPCodecFull
|
||||
)
|
||||
|
||||
type tcpTransportProbe struct {
|
||||
raw net.Conn
|
||||
prefix [8]byte
|
||||
prefixLen int
|
||||
mode tcpWireMode
|
||||
codec detectedTCPCodec
|
||||
}
|
||||
|
||||
// detectTCPTransport reads at most the first eight bytes once per physical TCP
|
||||
// connection. Telegram's obfuscated2 nonce generation explicitly excludes all
|
||||
// plaintext codec tags and requires bytes 4..8 to be non-zero, while the first
|
||||
// Full frame has transport sequence number zero. Those disjoint invariants let
|
||||
// one public port admit both wire modes without probabilistic guessing.
|
||||
func detectTCPTransport(raw net.Conn) (tcpTransportProbe, error) {
|
||||
probe := tcpTransportProbe{raw: raw}
|
||||
if _, err := io.ReadFull(raw, probe.prefix[:1]); err != nil {
|
||||
return probe, fmt.Errorf("read first transport byte: %w", err)
|
||||
}
|
||||
probe.prefixLen = 1
|
||||
if probe.prefix[0] == codec.AbridgedClientStart[0] {
|
||||
probe.mode = tcpWireModePlain
|
||||
probe.codec = detectedTCPCodecAbridged
|
||||
return probe, nil
|
||||
}
|
||||
|
||||
if _, err := io.ReadFull(raw, probe.prefix[1:4]); err != nil {
|
||||
return probe, fmt.Errorf("read transport prefix: %w", err)
|
||||
}
|
||||
probe.prefixLen = 4
|
||||
var firstFour [4]byte
|
||||
copy(firstFour[:], probe.prefix[:4])
|
||||
switch firstFour {
|
||||
case codec.IntermediateClientStart:
|
||||
probe.mode = tcpWireModePlain
|
||||
probe.codec = detectedTCPCodecIntermediate
|
||||
return probe, nil
|
||||
case codec.PaddedIntermediateClientStart:
|
||||
probe.mode = tcpWireModePlain
|
||||
probe.codec = detectedTCPCodecPaddedIntermediate
|
||||
return probe, nil
|
||||
}
|
||||
|
||||
first := binary.LittleEndian.Uint32(firstFour[:])
|
||||
if isHTTPHeaderPrefix(firstFour) || first == 0x02010316 {
|
||||
return probe, fmt.Errorf("%w: reserved prefix %x", errInvalidTCPTransportPrefix, firstFour)
|
||||
}
|
||||
|
||||
if _, err := io.ReadFull(raw, probe.prefix[4:8]); err != nil {
|
||||
return probe, fmt.Errorf("read transport discriminator: %w", err)
|
||||
}
|
||||
probe.prefixLen = 8
|
||||
if binary.LittleEndian.Uint32(probe.prefix[4:8]) == 0 {
|
||||
length := binary.LittleEndian.Uint32(probe.prefix[:4])
|
||||
if length < 3*4 || length > maxTransportMessageSize || length%4 != 0 {
|
||||
return probe, fmt.Errorf("%w: invalid full header length %d", errInvalidTCPTransportPrefix, length)
|
||||
}
|
||||
probe.mode = tcpWireModePlain
|
||||
probe.codec = detectedTCPCodecFull
|
||||
return probe, nil
|
||||
}
|
||||
|
||||
probe.mode = tcpWireModeObfuscated
|
||||
return probe, nil
|
||||
}
|
||||
|
||||
func (p tcpTransportProbe) replayConn() net.Conn {
|
||||
return newReplayNetConn(p.raw, p.prefix[:p.prefixLen])
|
||||
}
|
||||
|
||||
func (p tcpTransportProbe) plainFrameConn() net.Conn {
|
||||
if p.codec == detectedTCPCodecFull {
|
||||
// Full has no standalone codec tag: the first eight bytes are already the
|
||||
// length and sequence number of its first frame and must be replayed.
|
||||
return p.replayConn()
|
||||
}
|
||||
// Abridged/intermediate/padded-intermediate tags are client-only connection
|
||||
// headers. They were consumed by detection and are not part of the first frame.
|
||||
return p.raw
|
||||
}
|
||||
|
||||
type replayNetConn struct {
|
||||
net.Conn
|
||||
prefix [8]byte
|
||||
n uint8
|
||||
offset uint8
|
||||
}
|
||||
|
||||
func newReplayNetConn(conn net.Conn, prefix []byte) *replayNetConn {
|
||||
if len(prefix) > 8 {
|
||||
panic("mtprotoedge: replay prefix exceeds fixed transport discriminator")
|
||||
}
|
||||
replayed := &replayNetConn{Conn: conn, n: uint8(len(prefix))}
|
||||
copy(replayed.prefix[:], prefix)
|
||||
return replayed
|
||||
}
|
||||
|
||||
func (c *replayNetConn) Read(p []byte) (int, error) {
|
||||
if c.offset < c.n {
|
||||
n := copy(p, c.prefix[c.offset:c.n])
|
||||
c.offset += uint8(n)
|
||||
return n, nil
|
||||
}
|
||||
return c.Conn.Read(p)
|
||||
}
|
||||
|
||||
type serverObfuscatedConn struct {
|
||||
net.Conn
|
||||
rw io.ReadWriter
|
||||
}
|
||||
|
||||
func (c *serverObfuscatedConn) Read(p []byte) (int, error) {
|
||||
return c.rw.Read(p)
|
||||
}
|
||||
|
||||
func (c *serverObfuscatedConn) Write(p []byte) (int, error) {
|
||||
return c.rw.Write(p)
|
||||
}
|
||||
|
||||
func detectedCompatCodec(kind detectedTCPCodec) (transport.Codec, error) {
|
||||
switch kind {
|
||||
case detectedTCPCodecAbridged:
|
||||
return &quickAckAbridgedCodec{}, nil
|
||||
case detectedTCPCodecIntermediate:
|
||||
return &quickAckIntermediateCodec{}, nil
|
||||
case detectedTCPCodecPaddedIntermediate:
|
||||
return &quickAckPaddedIntermediateCodec{}, nil
|
||||
case detectedTCPCodecFull:
|
||||
return transport.Full.Codec(), nil
|
||||
default:
|
||||
return nil, fmt.Errorf("%w: unknown detected codec %d", errInvalidTCPTransportPrefix, kind)
|
||||
}
|
||||
}
|
||||
|
||||
func detectObfuscatedProtocol(tag [4]byte) (detectedTCPCodec, int, error) {
|
||||
switch tag {
|
||||
case (codec.Abridged{}).ObfuscatedTag():
|
||||
return detectedTCPCodecAbridged, 1, nil
|
||||
case codec.IntermediateClientStart:
|
||||
return detectedTCPCodecIntermediate, 4, nil
|
||||
case codec.PaddedIntermediateClientStart:
|
||||
return detectedTCPCodecPaddedIntermediate, 4, nil
|
||||
default:
|
||||
return detectedTCPCodecUnknown, 0, fmt.Errorf("%w: %x", errInvalidObfuscatedProtocol, tag)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) promoteMixedTCP(raw net.Conn) (transport.Conn, string, error) {
|
||||
probe, err := detectTCPTransport(raw)
|
||||
if err != nil {
|
||||
return nil, "tcp_auto", err
|
||||
}
|
||||
if probe.mode == tcpWireModePlain {
|
||||
if s.codec != nil {
|
||||
conn, err := newCompatTransportConn(s.codec, probe.replayConn(), s.frameBudget)
|
||||
return conn, "tcp", err
|
||||
}
|
||||
connCodec, err := detectedCompatCodec(probe.codec)
|
||||
if err != nil {
|
||||
return nil, "tcp", err
|
||||
}
|
||||
conn, err := newCompatTransportConnWithCodec(
|
||||
probe.plainFrameConn(),
|
||||
connCodec,
|
||||
s.frameBudget,
|
||||
false,
|
||||
)
|
||||
return conn, "tcp", err
|
||||
}
|
||||
|
||||
replayed := probe.replayConn()
|
||||
rw, metadata, err := obfuscated2.Accept(replayed, nil)
|
||||
if err != nil {
|
||||
return nil, "obfuscated_tcp", fmt.Errorf("accept obfuscated2: %w", err)
|
||||
}
|
||||
obfuscated := &serverObfuscatedConn{Conn: replayed, rw: rw}
|
||||
detected, tagLen, err := detectObfuscatedProtocol(metadata.Protocol)
|
||||
if err != nil {
|
||||
return nil, "obfuscated_tcp", err
|
||||
}
|
||||
if s.codec != nil {
|
||||
conn, err := newCompatTransportConn(
|
||||
s.codec,
|
||||
newReplayNetConn(obfuscated, metadata.Protocol[:tagLen]),
|
||||
s.frameBudget,
|
||||
)
|
||||
return conn, "obfuscated_tcp", err
|
||||
}
|
||||
connCodec, err := detectedCompatCodec(detected)
|
||||
if err != nil {
|
||||
return nil, "obfuscated_tcp", err
|
||||
}
|
||||
conn, err := newCompatTransportConnWithCodec(
|
||||
obfuscated,
|
||||
connCodec,
|
||||
s.frameBudget,
|
||||
false,
|
||||
)
|
||||
return conn, "obfuscated_tcp", err
|
||||
}
|
||||
458
internal/mtprotoedge/transport_detection_test.go
Normal file
458
internal/mtprotoedge/transport_detection_test.go
Normal file
|
|
@ -0,0 +1,458 @@
|
|||
package mtprotoedge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap/zaptest"
|
||||
|
||||
"github.com/iamxvbaba/td/bin"
|
||||
"github.com/iamxvbaba/td/mtproxy"
|
||||
"github.com/iamxvbaba/td/mtproxy/obfuscator"
|
||||
"github.com/iamxvbaba/td/proto/codec"
|
||||
"github.com/iamxvbaba/td/transport"
|
||||
)
|
||||
|
||||
func TestDetectTCPTransportFragmentedPrefixes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
prefix []byte
|
||||
mode tcpWireMode
|
||||
codec detectedTCPCodec
|
||||
}{
|
||||
{name: "abridged", prefix: []byte{0xef}, mode: tcpWireModePlain, codec: detectedTCPCodecAbridged},
|
||||
{name: "intermediate", prefix: []byte{0xee, 0xee, 0xee, 0xee}, mode: tcpWireModePlain, codec: detectedTCPCodecIntermediate},
|
||||
{name: "padded", prefix: []byte{0xdd, 0xdd, 0xdd, 0xdd}, mode: tcpWireModePlain, codec: detectedTCPCodecPaddedIntermediate},
|
||||
{name: "full", prefix: []byte{12, 0, 0, 0, 0, 0, 0, 0}, mode: tcpWireModePlain, codec: detectedTCPCodecFull},
|
||||
{name: "obfuscated", prefix: []byte{1, 2, 3, 4, 5, 6, 7, 8}, mode: tcpWireModeObfuscated, codec: detectedTCPCodecUnknown},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
server, client := net.Pipe()
|
||||
defer func() { _ = server.Close() }()
|
||||
defer func() { _ = client.Close() }()
|
||||
writeErr := make(chan error, 1)
|
||||
go func() {
|
||||
for _, b := range tt.prefix {
|
||||
if _, err := client.Write([]byte{b}); err != nil {
|
||||
writeErr <- err
|
||||
return
|
||||
}
|
||||
}
|
||||
writeErr <- nil
|
||||
}()
|
||||
|
||||
probe, err := detectTCPTransport(server)
|
||||
if err != nil {
|
||||
t.Fatalf("detect: %v", err)
|
||||
}
|
||||
if probe.mode != tt.mode || probe.codec != tt.codec || probe.prefixLen != len(tt.prefix) {
|
||||
t.Fatalf("probe = mode:%d codec:%d prefix:%d, want %d/%d/%d",
|
||||
probe.mode, probe.codec, probe.prefixLen, tt.mode, tt.codec, len(tt.prefix))
|
||||
}
|
||||
if err := <-writeErr; err != nil {
|
||||
t.Fatalf("write prefix: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetectTCPTransportRejectsReservedAndInvalidFullPrefixes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
prefix []byte
|
||||
}{
|
||||
{name: "http", prefix: []byte("GET ")},
|
||||
{name: "reserved", prefix: []byte{0x16, 0x03, 0x01, 0x02}},
|
||||
{name: "full_too_short", prefix: []byte{8, 0, 0, 0, 0, 0, 0, 0}},
|
||||
{name: "full_unaligned", prefix: []byte{13, 0, 0, 0, 0, 0, 0, 0}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
server, client := net.Pipe()
|
||||
defer func() { _ = server.Close() }()
|
||||
defer func() { _ = client.Close() }()
|
||||
go func() {
|
||||
_, _ = client.Write(tt.prefix)
|
||||
}()
|
||||
_, err := detectTCPTransport(server)
|
||||
if !errors.Is(err, errInvalidTCPTransportPrefix) {
|
||||
t.Fatalf("detect error = %v, want invalid prefix", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerMixedTCPAcceptsEveryPlainCodec(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
protocol transport.Protocol
|
||||
}{
|
||||
{name: "abridged", protocol: transport.Abridged},
|
||||
{name: "intermediate", protocol: transport.Intermediate},
|
||||
{name: "padded_intermediate", protocol: transport.PaddedIntermediate},
|
||||
{name: "full", protocol: transport.Full},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
addr, frames := startTransportFrameServer(t, Options{
|
||||
Logger: zaptest.NewLogger(t),
|
||||
ObfuscatedTCP: true,
|
||||
})
|
||||
raw, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
conn, err := tt.protocol.Handshake(raw)
|
||||
if err != nil {
|
||||
_ = raw.Close()
|
||||
t.Fatalf("transport handshake: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
|
||||
var payload bin.Buffer
|
||||
payload.PutInt32(0x12345678)
|
||||
payload.PutInt32(0x0badf00d)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
if err := conn.Send(ctx, &payload); err != nil {
|
||||
t.Fatalf("send: %v", err)
|
||||
}
|
||||
select {
|
||||
case n := <-frames:
|
||||
if n < payload.Len() || n > payload.Len()+15 {
|
||||
t.Fatalf("frame len = %d, want payload %d plus at most padded-intermediate padding", n, payload.Len())
|
||||
}
|
||||
case <-ctx.Done():
|
||||
t.Fatal("server did not receive plain frame in mixed mode")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerMixedTCPAcceptsObfuscatedCodecs(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
tag [4]byte
|
||||
newCodec func() transport.Codec
|
||||
}{
|
||||
{
|
||||
name: "abridged",
|
||||
tag: (codec.Abridged{}).ObfuscatedTag(),
|
||||
newCodec: func() transport.Codec {
|
||||
return codec.NoHeader{Codec: codec.Abridged{}}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "intermediate",
|
||||
tag: codec.IntermediateClientStart,
|
||||
newCodec: func() transport.Codec {
|
||||
return codec.NoHeader{Codec: codec.Intermediate{}}
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "padded_intermediate",
|
||||
tag: codec.PaddedIntermediateClientStart,
|
||||
newCodec: func() transport.Codec {
|
||||
return codec.NoHeader{Codec: codec.PaddedIntermediate{}}
|
||||
},
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
addr, frames := startTransportFrameServer(t, Options{
|
||||
Logger: zaptest.NewLogger(t),
|
||||
ObfuscatedTCP: true,
|
||||
})
|
||||
raw, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
obfuscated := obfuscator.Obfuscated2(rand.Reader, raw)
|
||||
if err := obfuscated.Handshake(tt.tag, 2, mtproxy.Secret{}); err != nil {
|
||||
_ = raw.Close()
|
||||
t.Fatalf("obfuscated handshake: %v", err)
|
||||
}
|
||||
conn, err := transport.NewProtocol(tt.newCodec).Handshake(obfuscated)
|
||||
if err != nil {
|
||||
_ = raw.Close()
|
||||
t.Fatalf("transport handshake: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
|
||||
var payload bin.Buffer
|
||||
payload.PutInt32(0x12345678)
|
||||
payload.PutInt32(0x0badf00d)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
if err := conn.Send(ctx, &payload); err != nil {
|
||||
t.Fatalf("send: %v", err)
|
||||
}
|
||||
select {
|
||||
case n := <-frames:
|
||||
if n < payload.Len() || n > payload.Len()+15 {
|
||||
t.Fatalf("frame len = %d, want payload %d plus at most padding", n, payload.Len())
|
||||
}
|
||||
case <-ctx.Done():
|
||||
t.Fatal("server did not receive obfuscated frame")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerMixedTCPPlainIntermediateKeyExchange(t *testing.T) {
|
||||
addr, pub, _ := startTestServer(t, Options{DC: 2, ObfuscatedTCP: true})
|
||||
conn, _, _ := dialHandshake(t, addr, 2, pub)
|
||||
_ = conn.Close()
|
||||
}
|
||||
|
||||
func TestServerMixedTCPResponseUsesDetectedWireMode(t *testing.T) {
|
||||
plainDial := func(protocol transport.Protocol) func(*testing.T, string) transport.Conn {
|
||||
return func(t *testing.T, addr string) transport.Conn {
|
||||
t.Helper()
|
||||
raw, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
conn, err := protocol.Handshake(raw)
|
||||
if err != nil {
|
||||
_ = raw.Close()
|
||||
t.Fatalf("transport handshake: %v", err)
|
||||
}
|
||||
return conn
|
||||
}
|
||||
}
|
||||
obfuscatedDial := func(
|
||||
tag [4]byte,
|
||||
newCodec func() transport.Codec,
|
||||
) func(*testing.T, string) transport.Conn {
|
||||
return func(t *testing.T, addr string) transport.Conn {
|
||||
t.Helper()
|
||||
raw, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
obfuscated := obfuscator.Obfuscated2(rand.Reader, raw)
|
||||
if err := obfuscated.Handshake(tag, 2, mtproxy.Secret{}); err != nil {
|
||||
_ = raw.Close()
|
||||
t.Fatalf("obfuscated handshake: %v", err)
|
||||
}
|
||||
conn, err := transport.NewProtocol(newCodec).Handshake(obfuscated)
|
||||
if err != nil {
|
||||
_ = raw.Close()
|
||||
t.Fatalf("transport handshake: %v", err)
|
||||
}
|
||||
return conn
|
||||
}
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
dial func(t *testing.T, addr string) transport.Conn
|
||||
}{
|
||||
{
|
||||
name: "plain_abridged",
|
||||
dial: plainDial(transport.Abridged),
|
||||
},
|
||||
{
|
||||
name: "plain_intermediate",
|
||||
dial: plainDial(transport.Intermediate),
|
||||
},
|
||||
{
|
||||
name: "plain_padded_intermediate",
|
||||
dial: plainDial(transport.PaddedIntermediate),
|
||||
},
|
||||
{
|
||||
name: "plain_full",
|
||||
dial: plainDial(transport.Full),
|
||||
},
|
||||
{
|
||||
name: "obfuscated_abridged",
|
||||
dial: obfuscatedDial(
|
||||
(codec.Abridged{}).ObfuscatedTag(),
|
||||
func() transport.Codec {
|
||||
return codec.NoHeader{Codec: codec.Abridged{}}
|
||||
},
|
||||
),
|
||||
},
|
||||
{
|
||||
name: "obfuscated_intermediate",
|
||||
dial: obfuscatedDial(
|
||||
codec.IntermediateClientStart,
|
||||
func() transport.Codec {
|
||||
return codec.NoHeader{Codec: codec.Intermediate{}}
|
||||
},
|
||||
),
|
||||
},
|
||||
{
|
||||
name: "obfuscated_padded_intermediate",
|
||||
dial: obfuscatedDial(
|
||||
codec.PaddedIntermediateClientStart,
|
||||
func() transport.Codec {
|
||||
return codec.NoHeader{Codec: codec.PaddedIntermediate{}}
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
addr, _, _ := startTestServer(t, Options{DC: 2, ObfuscatedTCP: true})
|
||||
conn := tt.dial(t, addr)
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
|
||||
var payload bin.Buffer
|
||||
payload.PutLong(0x1020304050607080) // deliberately unknown auth_key_id
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
if err := conn.Send(ctx, &payload); err != nil {
|
||||
t.Fatalf("send: %v", err)
|
||||
}
|
||||
var response bin.Buffer
|
||||
err := conn.Recv(ctx, &response)
|
||||
var protocolErr *codec.ProtocolErr
|
||||
if !errors.As(err, &protocolErr) || protocolErr.Code != codec.CodeAuthKeyNotFound {
|
||||
t.Fatalf("response error = %T %v, want transport -404 in the detected wire mode", err, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerMixedTCPRejectsInvalidPrefixesWithoutObfuscationWait(t *testing.T) {
|
||||
addr, _, _ := startTestServer(t, Options{
|
||||
ObfuscatedTCP: true,
|
||||
HandshakeIdleTimeout: 5 * time.Second,
|
||||
})
|
||||
raw, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
defer func() { _ = raw.Close() }()
|
||||
if _, err := raw.Write([]byte{8, 0, 0, 0, 0, 0, 0, 0}); err != nil {
|
||||
t.Fatalf("write invalid full prefix: %v", err)
|
||||
}
|
||||
if err := raw.SetReadDeadline(time.Now().Add(time.Second)); err != nil {
|
||||
t.Fatalf("set read deadline: %v", err)
|
||||
}
|
||||
_, err = raw.Read(make([]byte, 1))
|
||||
if err == nil {
|
||||
t.Fatal("invalid transport prefix unexpectedly kept connection open")
|
||||
}
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) && netErr.Timeout() {
|
||||
t.Fatalf("invalid prefix waited for obfuscation timeout instead of failing fast: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerMixedTCPRejectsUnknownObfuscatedProtocolTag(t *testing.T) {
|
||||
addr, _, _ := startTestServer(t, Options{ObfuscatedTCP: true})
|
||||
raw, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
defer func() { _ = raw.Close() }()
|
||||
obfuscated := obfuscator.Obfuscated2(rand.Reader, raw)
|
||||
if err := obfuscated.Handshake([4]byte{1, 2, 3, 4}, 2, mtproxy.Secret{}); err != nil {
|
||||
t.Fatalf("obfuscated handshake: %v", err)
|
||||
}
|
||||
if err := raw.SetReadDeadline(time.Now().Add(time.Second)); err != nil {
|
||||
t.Fatalf("set read deadline: %v", err)
|
||||
}
|
||||
_, err = obfuscated.Read(make([]byte, 1))
|
||||
if err == nil {
|
||||
t.Fatal("unknown obfuscated protocol tag unexpectedly kept connection open")
|
||||
}
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) && netErr.Timeout() {
|
||||
t.Fatalf("unknown obfuscated protocol tag did not fail closed: %v", err)
|
||||
}
|
||||
if !errors.Is(err, io.EOF) {
|
||||
// Windows may report a reset instead of EOF; any non-timeout terminal
|
||||
// network error is the same fail-closed outcome.
|
||||
t.Logf("terminal read error after invalid obfuscated tag: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func startTransportFrameServer(t *testing.T, opts Options) (string, <-chan int) {
|
||||
t.Helper()
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
frames := make(chan int, 1)
|
||||
srv := New(opts)
|
||||
srv.onFrame = func(n int) {
|
||||
select {
|
||||
case frames <- n:
|
||||
default:
|
||||
}
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
serveErr := make(chan error, 1)
|
||||
go func() { serveErr <- srv.Serve(ctx, ln) }()
|
||||
t.Cleanup(func() {
|
||||
cancel()
|
||||
select {
|
||||
case err := <-serveErr:
|
||||
if err != nil {
|
||||
t.Errorf("serve: %v", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Error("server did not stop")
|
||||
}
|
||||
})
|
||||
return ln.Addr().String(), frames
|
||||
}
|
||||
|
||||
func BenchmarkDetectTCPTransport(b *testing.B) {
|
||||
tests := []struct {
|
||||
name string
|
||||
prefix []byte
|
||||
}{
|
||||
{name: "plain_abridged", prefix: []byte{0xef}},
|
||||
{name: "plain_intermediate", prefix: []byte{0xee, 0xee, 0xee, 0xee}},
|
||||
{name: "plain_full", prefix: []byte{12, 0, 0, 0, 0, 0, 0, 0}},
|
||||
{name: "obfuscated", prefix: []byte{1, 2, 3, 4, 5, 6, 7, 8}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
b.Run(tt.name, func(b *testing.B) {
|
||||
conn := &transportProbeBenchmarkConn{payload: tt.prefix}
|
||||
b.ReportAllocs()
|
||||
for i := 0; i < b.N; i++ {
|
||||
conn.offset = 0
|
||||
if _, err := detectTCPTransport(conn); err != nil {
|
||||
b.Fatal(err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type transportProbeBenchmarkConn struct {
|
||||
payload []byte
|
||||
offset int
|
||||
}
|
||||
|
||||
func (c *transportProbeBenchmarkConn) Read(p []byte) (int, error) {
|
||||
if c.offset >= len(c.payload) {
|
||||
return 0, io.EOF
|
||||
}
|
||||
n := copy(p, c.payload[c.offset:])
|
||||
c.offset += n
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func (*transportProbeBenchmarkConn) Write(p []byte) (int, error) { return len(p), nil }
|
||||
func (*transportProbeBenchmarkConn) Close() error { return nil }
|
||||
func (*transportProbeBenchmarkConn) LocalAddr() net.Addr { return nil }
|
||||
func (*transportProbeBenchmarkConn) RemoteAddr() net.Addr { return nil }
|
||||
func (*transportProbeBenchmarkConn) SetDeadline(time.Time) error { return nil }
|
||||
func (*transportProbeBenchmarkConn) SetReadDeadline(time.Time) error {
|
||||
return nil
|
||||
}
|
||||
func (*transportProbeBenchmarkConn) SetWriteDeadline(time.Time) error {
|
||||
return nil
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue