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

View file

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

View file

@ -89,7 +89,7 @@ func connLocal(conn net.Conn) string {
func intakeTransport(obfuscated bool) string { func intakeTransport(obfuscated bool) string {
if obfuscated { if obfuscated {
return "obfuscated_tcp" return "tcp_auto"
} }
return "tcp" 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) { func newFrameBudgetTestTransport(packet []byte, c transport.Codec, budget *inboundFrameBudget) (*compatTransportConn, *frameBudgetTestConn) {
raw := newFrameBudgetTestConn(packet) 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) { func TestInboundFrameBudgetSupportsBuiltInCodecs(t *testing.T) {

View file

@ -1,7 +1,6 @@
package mtprotoedge package mtprotoedge
import ( import (
"bytes"
"context" "context"
"errors" "errors"
"io" "io"
@ -195,10 +194,9 @@ func (m *samePortMux) dispatch(ctx context.Context, conn net.Conn) {
return return
} }
wrapped := &prefixedNetConn{ // Fixed eight-byte replay storage avoids bytes.Reader + MultiReader allocations
Conn: conn, // on every accepted same-port connection.
reader: io.MultiReader(bytes.NewReader(header[:]), conn), wrapped := newReplayNetConn(conn, header[:])
}
target := m.tcp target := m.tcp
transport := "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 把分流后的连接投递进来,下游 // samePortMuxListener 是一个内存 listener:dispatch 把分流后的连接投递进来,下游
// (serveMixed 的 accept 循环 / http.Server) 从这里 Accept。 // (serveMixed 的 accept 循环 / http.Server) 从这里 Accept。
type samePortMuxListener struct { type samePortMuxListener struct {

View file

@ -221,12 +221,14 @@ type Options struct {
// codec 必须是 gotd 内置四种 codec(可包 NoHeader),或实现 InboundFrameBudgetedCodec; // codec 必须是 gotd 内置四种 codec(可包 NoHeader),或实现 InboundFrameBudgetedCodec;
// 无法在 payload 分配前预检长度的 codec 会 fail-closed。 // 无法在 payload 分配前预检长度的 codec 会 fail-closed。
Codec func() transport.Codec Codec func() transport.Codec
// ObfuscatedTCP 先按 MTProto TCP obfuscation 解包,再自动探测 codec。 // ObfuscatedTCP 允许裸 TCP 使用 MTProto transport obfuscation。开启时按每条
// Telegram Desktop 的 tcpo_only endpoint 会走这个 64 字节前缀流程。 // 物理连接的首 1/4/8 字节自动区分明文 transport 与 64-byte obfuscated2,
// 随后把 wire mode + codec 冻结到同一条双向连接;Telegram Desktop 的
// tcpo_only 与未开启混淆的第三方客户端可共用同一端口。
ObfuscatedTCP bool ObfuscatedTCP bool
// WebSocket 在同一个 listener 上接受 MTProto over WebSocket(/apiws*)。 // WebSocket 在同一个 listener 上接受 MTProto over WebSocket(/apiws*)。
// 开启后仅在连接建立时读取前 4 字节做 HTTP/TCP 分流;MTProto TCP // 开启后仅在连接建立时读取前 4 字节做 HTTP/TCP 分流;MTProto TCP
// 后续仍走原 ObfuscatedTCP + codec 热路径。 // 后续仍走原 TCP wire-mode + codec 探测路径。
WebSocket bool WebSocket bool
// WebSocketAllowedOrigins 是允许浏览器发起 WebSocket upgrade 的页面 origin。 // WebSocketAllowedOrigins 是允许浏览器发起 WebSocket upgrade 的页面 origin。
// 空列表表示只接受无 Origin 的非浏览器客户端;"*" 表示允许所有来源(仅调试)。 // 空列表表示只接受无 Origin 的非浏览器客户端;"*" 表示允许所有来源(仅调试)。
@ -644,7 +646,11 @@ func (s *Server) serveTCP(ctx context.Context, ln net.Listener) error {
ctx, cancel := context.WithCancel(ctx) ctx, cancel := context.WithCancel(ctx)
defer cancel() 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") defer s.log.Info("Stopped")
errCh := make(chan error, 1) errCh := make(chan error, 1)
go func() { go func() {
@ -680,7 +686,7 @@ func (s *Server) serveMixed(ctx context.Context, ln net.Listener) error {
s.log.Info("Serving", s.log.Info("Serving",
zap.String("addr", ln.Addr().String()), zap.String("addr", ln.Addr().String()),
zap.Int("dc", s.dc), zap.Int("dc", s.dc),
zap.Bool("obfuscated_tcp", s.obfuscated), zap.String("tcp_transport_mode", intakeTransport(s.obfuscated)),
zap.Bool("websocket", true), zap.Bool("websocket", true),
zap.Strings("websocket_origins", s.websocketOrigins), zap.Strings("websocket_origins", s.websocketOrigins),
) )
@ -705,7 +711,7 @@ func (s *Server) serveMixed(ctx context.Context, ln net.Listener) error {
defer wg.Done() defer wg.Done()
errCh <- mux.Serve(ctx) errCh <- mux.Serve(ctx)
}() }()
// 裸 MTProto TCP:每条连接在自己的 goroutine 里完成去混淆 + codec 探测。 // 裸 MTProto TCP:每条连接在自己的 goroutine 里完成 wire mode + codec 探测。
go func() { go func() {
defer wg.Done() defer wg.Done()
errCh <- s.acceptLoop(ctx, mux.TCP(), s.obfuscated) 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 return firstErr
} }
// acceptLoop 接受裸连接,并为每条连接单独起 goroutine 完成「去混淆 + codec 探测 + // acceptLoop 接受裸连接,并为每条连接单独起 goroutine 完成「wire mode + codec 探测 +
// serveConn」。探测在 accept 循环之外、带握手超时进行——慢/半开/坏 init 的客户端只占用 // serveConn」。探测在 accept 循环之外、带握手超时进行——慢/半开/坏 init 的客户端只占用
// 自己的 goroutine,绝不阻塞其他连接的接入;单条连接的握手失败也只关闭该连接,不会拖垮 // 自己的 goroutine,绝不阻塞其他连接的接入;单条连接的握手失败也只关闭该连接,不会拖垮
// 整个监听循环。obfuscated 为 true 时先走 obfuscated2 去混淆(裸 MTProto TCP);WebSocket // 整个监听循环。obfuscated 为 true 时自动区分 plain 与 obfuscated2(裸 MTProto TCP);
// 连接传 false(gotd 升级处理器已完成去混淆)。 // WebSocket 连接传 false(gotd 升级处理器已完成去混淆)。
func (s *Server) acceptLoop(ctx context.Context, ln net.Listener, obfuscated bool) error { func (s *Server) acceptLoop(ctx context.Context, ln net.Listener, obfuscated bool) error {
return s.acceptLoopTransport(ctx, ln, obfuscated, intakeTransport(obfuscated)) 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) { func (s *Server) serveDetectedConn(ctx context.Context, raw net.Conn, obfuscated bool, transportName string) {
started := time.Now() started := time.Now()
remote, local := connRemote(raw), connLocal(raw) remote, local := connRemote(raw), connLocal(raw)
// 握手读超时只覆盖去混淆 + codec 探测这一小段;用真实墙钟时间(SetReadDeadline 语义), // 握手读超时只覆盖 wire-mode + codec 探测这一小段;用真实墙钟时间(SetReadDeadline 语义),
// 不走可能被测试注入的逻辑 clock。 // 不走可能被测试注入的逻辑 clock。
if err := raw.SetReadDeadline(time.Now().Add(s.handshakeTimeout)); err != nil { if err := raw.SetReadDeadline(time.Now().Add(s.handshakeTimeout)); err != nil {
_ = raw.Close() _ = 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) close(promoted)
if err != nil { if err != nil {
outcome := "error" outcome := "error"
@ -867,15 +876,15 @@ func (s *Server) serveDetectedConn(ctx context.Context, raw net.Conn, obfuscated
} }
} }
// promoteConn 复用与 listener 组合完全一致的「obfuscated2 去混淆 + codec 探测」管线,但针对 // promoteConn 针对单条连接执行一次 transport 提升。obfuscated=true 表示裸 TCP
// 单条连接,使其可在 accept 循环之外执行。obfuscated 对 WebSocket 连接必须为 false(gotd // 允许混淆并自动区分 plain/obfuscated2;WebSocket 必须传 false,因为 gotd upgrade
// 升级处理器已剥离 obfuscated2 并补回 codec tag)。 // handler 已剥离 obfuscated2 并补回 codec tag。
func (s *Server) promoteConn(raw net.Conn, obfuscated bool) (transport.Conn, error) { func (s *Server) promoteConn(raw net.Conn, obfuscated bool) (transport.Conn, string, error) {
var ln net.Listener = newSingleConnListener(raw)
if obfuscated { 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 分流。 // 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") ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil { if err != nil {
t.Fatalf("listen: %v", err) 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 := New(Options{Logger: zaptest.NewLogger(t), ObfuscatedTCP: true, WebSocket: true})
srv.onFrame = func(n int) { srv.onFrame = func(n int) {
select { select {
@ -302,6 +302,27 @@ func TestServerSamePortWebSocketAndObfuscatedTCP(t *testing.T) {
expectFrameLen(t, frames, tcpPayload.Len()) expectFrameLen(t, frames, tcpPayload.Len())
_ = tcpConn.Close() _ = 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() cancel()
select { select {
case err := <-serveErr: case err := <-serveErr:

View file

@ -89,8 +89,7 @@ func newCompatTransportListener(codec func() transport.Codec, listener net.Liste
} }
// singleConnListener 是一个只产出一条「已接受」连接、随后阻塞到关闭的 net.Listener。 // singleConnListener 是一个只产出一条「已接受」连接、随后阻塞到关闭的 net.Listener。
// 它让单条裸连接可以走 listener 形态的去混淆/codec 管线(ObfuscatedListener + // 生产接入已直接提升单连接;这个适配器仅保留给仍需 listener 形态的测试/调用方。
// compatTransportListener),从而把这部分阻塞读取从 accept 循环挪到每连接 goroutine。
type singleConnListener struct { type singleConnListener struct {
addr net.Addr addr net.Addr
ch chan net.Conn ch chan net.Conn
@ -131,38 +130,13 @@ func (l *compatTransportListener) Accept() (_ transport.Conn, rErr error) {
} }
}() }()
var ( promoted, err := newCompatTransportConn(l.codec, conn, l.budget)
connCodec transport.Codec if err != nil {
reader io.Reader = conn // Avoid returning a typed nil *compatTransportConn as a non-nil
) // transport.Conn interface on admission failure.
if l.codec != nil { return nil, err
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")
}
} }
return promoted, nil
return &compatTransportConn{
conn: wrappedCompatConn{
reader: reader,
Conn: conn,
},
codec: connCodec,
budget: l.budget,
transportPacketMessages: isTransportPacketMessageConn(conn),
}, nil
} }
func isTransportPacketMessageConn(conn net.Conn) bool { func isTransportPacketMessageConn(conn net.Conn) bool {
@ -187,10 +161,87 @@ func (w wrappedCompatConn) Read(p []byte) (int, error) {
return w.reader.Read(p) 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 { type compatTransportConn struct {
conn net.Conn conn net.Conn
codec transport.Codec codec transport.Codec
budget *inboundFrameBudget codecKind inboundFrameCodecKind
budgetedCodec InboundFrameBudgetedCodec
budget *inboundFrameBudget
transportPacketMessages bool transportPacketMessages bool
directMessageScratch []byte directMessageScratch []byte
@ -427,7 +478,10 @@ func (c *compatTransportConn) Close() error {
} }
func (c *compatTransportConn) readInboundFrame(b *bin.Buffer) 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 { if kind == inboundFrameCodecUnknown {
return errInboundFrameCodecUnsupported return errInboundFrameCodecUnsupported
} }
@ -447,11 +501,10 @@ func (c *compatTransportConn) readInboundFrame(b *bin.Buffer) error {
var err error var err error
if kind == inboundFrameCodecCustom { if kind == inboundFrameCodecCustom {
custom := unwrapInboundFrameBudgetedCodec(c.codec) if c.budgetedCodec == nil {
if custom == nil {
return errInboundFrameCodecUnsupported return errInboundFrameCodecUnsupported
} }
err = custom.ReadWithInboundFrameBudget(c.conn, b, reserve) err = c.budgetedCodec.ReadWithInboundFrameBudget(c.conn, b, reserve)
} else { } else {
preflight := &inboundFramePreflightReader{r: c.conn, kind: kind, reserve: reserve} preflight := &inboundFramePreflightReader{r: c.conn, kind: kind, reserve: reserve}
err = c.codec.Read(preflight, b) 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
}