fix: sync same-port transport detection

This commit is contained in:
iamxvbaba 2026-07-24 14:50:17 +08:00
parent 840ef6237c
commit 2289a31f46
10 changed files with 851 additions and 83 deletions

View file

@ -9,8 +9,8 @@
// go run ./cmd/bots/botcheck -token "<bot_id>:<secret>" # 仅登录自检
// go run ./cmd/bots/botcheck -token "<bot_id>:<secret>" -echo # 自检后持续 echo
//
// 连接生产 telesrvobfuscated TCP靠 DCOption.TCPObfuscatedOnly=true
// gotd dcs.Plain 据此自动走 MTProto TCP obfuscation。
// 以 obfuscated TCP 连接生产 telesrvserver 会逐连接自动区分 plain/obfuscated
// 此探针靠 DCOption.TCPObfuscatedOnly=true 让 gotd 客户端选择 MTProto TCP obfuscation。
package main
import (
@ -42,7 +42,7 @@ import (
)
// obfuscatedResolver 用标准无-secret MTProto TCP obfuscationobfuscated2连接
// 匹配 telesrv 生产 server 的 transport.ObfuscatedListenerobfuscated2.Accept(conn, nil)
// 匹配 telesrv 生产 server 自动检测后的 obfuscated2.Accept(conn, nil) 路径
// gotd 内置 dcs.Plain 的 obfuscated 路径走 MTProxy强制 secret不适用这里。
type obfuscatedResolver struct {
host string

View file

@ -64,8 +64,8 @@ import (
)
// obfuscatedResolver 用标准无-secret MTProto TCP obfuscationobfuscated2连接 telesrv
// 匹配生产 server 的 transport.ObfuscatedListener。gotd 内置 dcs.Plain 的 obfuscated 路径
// 走 MTProxy强制 secret不适用这里所以自定义一个 Resolver。
// 匹配生产 server 自动检测后的 obfuscated2 路径。gotd 内置 dcs.Plain 的
// TCPObfuscatedOnly 路径走 MTProxy强制 secret不适用这里所以自定义 Resolver。
type obfuscatedResolver struct {
host string
port int

View file

@ -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"
}

View file

@ -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) {

View file

@ -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 是一个内存 listenerdispatch 把分流后的连接投递进来,下游
// (serveMixed 的 accept 循环 / http.Server) 从这里 Accept。
type samePortMuxListener struct {

View file

@ -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 TCPWebSocket
// 连接传 falsegotd 升级处理器已完成去混淆)。
// 整个监听循环。obfuscated 为 true 时自动区分 plain 与 obfuscated2裸 MTProto TCP
// WebSocket 连接传 falsegotd 升级处理器已完成去混淆)。
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 连接必须为 falsegotd
// 升级处理器已剥离 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/obfuscated2WebSocket 必须传 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 分流。

View file

@ -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:

View file

@ -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)

View 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
}

View 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
}