diff --git a/cmd/bots/botcheck/main.go b/cmd/bots/botcheck/main.go index 4de8b6c8..8c29e859 100644 --- a/cmd/bots/botcheck/main.go +++ b/cmd/bots/botcheck/main.go @@ -9,8 +9,8 @@ // go run ./cmd/bots/botcheck -token ":" # 仅登录自检 // go run ./cmd/bots/botcheck -token ":" -echo # 自检后持续 echo // -// 连接生产 telesrv(obfuscated TCP)靠 DCOption.TCPObfuscatedOnly=true, -// gotd dcs.Plain 据此自动走 MTProto TCP obfuscation。 +// 以 obfuscated TCP 连接生产 telesrv:server 会逐连接自动区分 plain/obfuscated; +// 此探针靠 DCOption.TCPObfuscatedOnly=true 让 gotd 客户端选择 MTProto TCP obfuscation。 package main import ( @@ -42,7 +42,7 @@ import ( ) // 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),不适用这里。 type obfuscatedResolver struct { host string diff --git a/cmd/bots/botdemo/main.go b/cmd/bots/botdemo/main.go index 383348f8..3c5cc39b 100644 --- a/cmd/bots/botdemo/main.go +++ b/cmd/bots/botdemo/main.go @@ -64,8 +64,8 @@ import ( ) // obfuscatedResolver 用标准无-secret MTProto TCP obfuscation(obfuscated2)连接 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 diff --git a/internal/mtprotoedge/connection_intake.go b/internal/mtprotoedge/connection_intake.go index e77add6e..c1ea3702 100644 --- a/internal/mtprotoedge/connection_intake.go +++ b/internal/mtprotoedge/connection_intake.go @@ -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" } diff --git a/internal/mtprotoedge/frame_budget_test.go b/internal/mtprotoedge/frame_budget_test.go index a0ef896c..34095e53 100644 --- a/internal/mtprotoedge/frame_budget_test.go +++ b/internal/mtprotoedge/frame_budget_test.go @@ -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) { diff --git a/internal/mtprotoedge/same_port_mux.go b/internal/mtprotoedge/same_port_mux.go index 4862ebf4..c51e474a 100644 --- a/internal/mtprotoedge/same_port_mux.go +++ b/internal/mtprotoedge/same_port_mux.go @@ -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 { diff --git a/internal/mtprotoedge/server.go b/internal/mtprotoedge/server.go index d3665d6b..d1e7de6a 100644 --- a/internal/mtprotoedge/server.go +++ b/internal/mtprotoedge/server.go @@ -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 分流。 diff --git a/internal/mtprotoedge/server_test.go b/internal/mtprotoedge/server_test.go index c4812b41..a80571b8 100644 --- a/internal/mtprotoedge/server_test.go +++ b/internal/mtprotoedge/server_test.go @@ -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: diff --git a/internal/mtprotoedge/transport_compat.go b/internal/mtprotoedge/transport_compat.go index 7b42229b..999a0d92 100644 --- a/internal/mtprotoedge/transport_compat.go +++ b/internal/mtprotoedge/transport_compat.go @@ -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) diff --git a/internal/mtprotoedge/transport_detection.go b/internal/mtprotoedge/transport_detection.go new file mode 100644 index 00000000..04e7f93a --- /dev/null +++ b/internal/mtprotoedge/transport_detection.go @@ -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 +} diff --git a/internal/mtprotoedge/transport_detection_test.go b/internal/mtprotoedge/transport_detection_test.go new file mode 100644 index 00000000..f61f3f15 --- /dev/null +++ b/internal/mtprotoedge/transport_detection_test.go @@ -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 +}