owpengram-server/internal/mtprotoedge/transport_detection_test.go
2026-07-24 14:50:17 +08:00

458 lines
13 KiB
Go

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
}