chore: refresh gramsrv public release
This commit is contained in:
parent
75cebe8dbf
commit
70b6820474
1274 changed files with 378751 additions and 59919 deletions
|
|
@ -1,9 +1,16 @@
|
|||
package mtprotoedge
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
|
|
@ -13,6 +20,7 @@ import (
|
|||
"github.com/gotd/td/mtproxy"
|
||||
"github.com/gotd/td/mtproxy/obfuscator"
|
||||
"github.com/gotd/td/proto/codec"
|
||||
"github.com/gotd/td/telegram/dcs"
|
||||
"github.com/gotd/td/transport"
|
||||
)
|
||||
|
||||
|
|
@ -223,3 +231,395 @@ func TestServerAcceptObfuscatedAbridgedQuickAckFrame(t *testing.T) {
|
|||
t.Fatal("server did not stop after ctx cancel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerSamePortWebSocketAndObfuscatedTCP(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)
|
||||
srv := New(Options{Logger: zaptest.NewLogger(t), ObfuscatedTCP: true, WebSocket: true})
|
||||
srv.onFrame = func(n int) {
|
||||
select {
|
||||
case frames <- n:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
serveErr := make(chan error, 1)
|
||||
go func() { serveErr <- srv.Serve(ctx, ln) }()
|
||||
|
||||
var wsPayload bin.Buffer
|
||||
wsPayload.PutInt32(0x11223344)
|
||||
wsPayload.PutInt32(0x55667788)
|
||||
|
||||
wsResolver := dcs.Websocket(dcs.WebsocketOptions{})
|
||||
wsConn, err := wsResolver.Primary(context.Background(), 2, dcs.List{
|
||||
Domains: map[int]string{
|
||||
2: "ws://" + ln.Addr().String() + "/apiws",
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("websocket dial: %v", err)
|
||||
}
|
||||
sendCtx, sc := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
if err := wsConn.Send(sendCtx, &wsPayload); err != nil {
|
||||
sc()
|
||||
t.Fatalf("websocket send: %v", err)
|
||||
}
|
||||
sc()
|
||||
expectFrameLen(t, frames, wsPayload.Len())
|
||||
_ = wsConn.Close()
|
||||
|
||||
raw, err := net.Dial("tcp", ln.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatalf("tcp dial: %v", err)
|
||||
}
|
||||
obfs := obfuscator.Obfuscated2(rand.Reader, raw)
|
||||
if err := obfs.Handshake((codec.Abridged{}).ObfuscatedTag(), 2, mtproxy.Secret{}); err != nil {
|
||||
t.Fatalf("tcp obfuscated handshake: %v", err)
|
||||
}
|
||||
tcpConn, err := transport.NewProtocol(func() transport.Codec {
|
||||
return transport.Abridged.CodecNoHeader()
|
||||
}).Handshake(obfs)
|
||||
if err != nil {
|
||||
t.Fatalf("tcp transport handshake: %v", err)
|
||||
}
|
||||
|
||||
var tcpPayload bin.Buffer
|
||||
tcpPayload.PutInt32(0x12345678)
|
||||
tcpPayload.PutInt32(0x0badf00d)
|
||||
sendCtx, sc = context.WithTimeout(context.Background(), 5*time.Second)
|
||||
if err := tcpConn.Send(sendCtx, &tcpPayload); err != nil {
|
||||
sc()
|
||||
t.Fatalf("tcp send: %v", err)
|
||||
}
|
||||
sc()
|
||||
expectFrameLen(t, frames, tcpPayload.Len())
|
||||
_ = tcpConn.Close()
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case err := <-serveErr:
|
||||
if err != nil {
|
||||
t.Fatalf("serve returned error: %v", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("server did not stop after ctx cancel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSamePortWebSocketTransportRoundTrip(t *testing.T) {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
mux := newSamePortMux(ln, 5*time.Second)
|
||||
wsLn, wsHandler := transport.WebsocketListener(ln.Addr())
|
||||
httpServer := &http.Server{
|
||||
Handler: websocketRouteHandler(wsHandler, []string{"http://localhost:1234"}),
|
||||
ReadHeaderTimeout: 5 * time.Second,
|
||||
}
|
||||
|
||||
serveErr := make(chan error, 2)
|
||||
go func() { serveErr <- mux.Serve(ctx) }()
|
||||
go func() {
|
||||
err := httpServer.Serve(mux.HTTP())
|
||||
if err != nil && !errors.Is(err, http.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
|
||||
serveErr <- err
|
||||
return
|
||||
}
|
||||
serveErr <- nil
|
||||
}()
|
||||
defer func() {
|
||||
cancel()
|
||||
_ = mux.Close()
|
||||
_ = httpServer.Close()
|
||||
_ = wsLn.Close()
|
||||
for i := 0; i < 2; i++ {
|
||||
select {
|
||||
case err := <-serveErr:
|
||||
if err != nil {
|
||||
t.Fatalf("serve returned error: %v", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("same-port websocket transport did not stop")
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
serverDone := make(chan error, 1)
|
||||
go func() {
|
||||
l := newCompatTransportListener(nil, wsLn)
|
||||
defer func() { _ = l.Close() }()
|
||||
|
||||
conn, err := l.Accept()
|
||||
if err != nil {
|
||||
serverDone <- err
|
||||
return
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
|
||||
recvCtx, rc := context.WithTimeout(ctx, 5*time.Second)
|
||||
defer rc()
|
||||
var got bin.Buffer
|
||||
if err := conn.Recv(recvCtx, &got); err != nil {
|
||||
serverDone <- err
|
||||
return
|
||||
}
|
||||
|
||||
var reply bin.Buffer
|
||||
reply.PutInt32(0x10203040)
|
||||
reply.PutInt32(0x50607080)
|
||||
sendCtx, sc := context.WithTimeout(ctx, 5*time.Second)
|
||||
defer sc()
|
||||
if err := conn.Send(sendCtx, &reply); err != nil {
|
||||
serverDone <- err
|
||||
return
|
||||
}
|
||||
serverDone <- nil
|
||||
}()
|
||||
|
||||
wsResolver := dcs.Websocket(dcs.WebsocketOptions{})
|
||||
wsConn, err := wsResolver.Primary(context.Background(), 2, dcs.List{
|
||||
Domains: map[int]string{
|
||||
2: "ws://" + ln.Addr().String() + "/apiws",
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("websocket dial: %v", err)
|
||||
}
|
||||
defer func() { _ = wsConn.Close() }()
|
||||
|
||||
var request bin.Buffer
|
||||
request.PutInt32(0x11223344)
|
||||
request.PutInt32(0x55667788)
|
||||
sendCtx, sc := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
if err := wsConn.Send(sendCtx, &request); err != nil {
|
||||
sc()
|
||||
t.Fatalf("websocket send: %v", err)
|
||||
}
|
||||
sc()
|
||||
|
||||
var want bin.Buffer
|
||||
want.PutInt32(0x10203040)
|
||||
want.PutInt32(0x50607080)
|
||||
recvCtx, rc := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
var got bin.Buffer
|
||||
if err := wsConn.Recv(recvCtx, &got); err != nil {
|
||||
rc()
|
||||
t.Fatalf("websocket recv: %v", err)
|
||||
}
|
||||
rc()
|
||||
if !bytes.Equal(got.Raw(), want.Raw()) {
|
||||
t.Fatalf("websocket recv = %x, want %x", got.Raw(), want.Raw())
|
||||
}
|
||||
|
||||
select {
|
||||
case err := <-serverDone:
|
||||
if err != nil {
|
||||
t.Fatalf("server transport: %v", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("server transport did not finish")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSamePortMuxIdleConnNotReapedBeforeHandshakeTimeout 回归:WebSocket 同端口复用的嗅探
|
||||
// 读超时(读首 4 字节做 HTTP/TCP 分流)必须对齐 HandshakeIdleTimeout,而不是旧的硬上限 5s。
|
||||
// 合法 MTProto 客户端(DrKLO)会预开「暖」连接、在有请求前并不立即发 obfuscated2 init;旧实现
|
||||
// 用 minDuration(5s, handshakeTimeout) 把嗅探压到 5s,比非 mux 路径(serveDetectedConn 用满
|
||||
// handshakeTimeout)激进 12 倍,会把这些暖连接在 5s 误杀,触发 DrKLO 6s 重连风暴 + EPOLLRDHUP
|
||||
// + 误判后端不健康回退外部 DNS(见 docs/client-compat-notes.md)。
|
||||
//
|
||||
// HandshakeIdleTimeout 必须 >5s 才能暴露旧的截断:连一条裸 TCP、不发任何字节,本地用 6s 读
|
||||
// deadline。新实现下连接在 8s 嗅探超时前一直存活,故本地 Read 因自身 deadline 超时(net timeout);
|
||||
// 旧实现下服务端 5s 即 FIN,本地 Read 会在 6s 前拿到 EOF/reset —— 据此判定回归。
|
||||
func TestSamePortMuxIdleConnNotReapedBeforeHandshakeTimeout(t *testing.T) {
|
||||
addr, _, _ := startTestServer(t, Options{
|
||||
WebSocket: true,
|
||||
ObfuscatedTCP: true,
|
||||
HandshakeIdleTimeout: 8 * time.Second, // 必须 >5s 才能区分「对齐 handshakeTimeout」与旧的 5s 截断
|
||||
})
|
||||
|
||||
raw, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
defer func() { _ = raw.Close() }()
|
||||
|
||||
// 不发任何字节,模拟客户端预开、暂未发首帧的暖连接。
|
||||
if err := raw.SetReadDeadline(time.Now().Add(6 * time.Second)); err != nil {
|
||||
t.Fatalf("set read deadline: %v", err)
|
||||
}
|
||||
n, err := raw.Read(make([]byte, 1))
|
||||
if err == nil {
|
||||
t.Fatalf("unexpected %d bytes on idle pre-handshake conn (server should send nothing)", n)
|
||||
}
|
||||
// 只有「本地读 deadline 超时」才说明连接在 6s 时仍存活(嗅探超时已对齐 8s)。
|
||||
// 任何由对端关闭导致的 EOF/reset 都意味着服务端在 handshakeTimeout 前过早回收了连接。
|
||||
var nerr net.Error
|
||||
if errors.As(err, &nerr) && nerr.Timeout() {
|
||||
return
|
||||
}
|
||||
t.Fatalf("server closed idle pre-handshake conn before handshake idle timeout (err=%v); "+
|
||||
"same-port mux sniff timeout must align with HandshakeIdleTimeout, not a 5s cap", err)
|
||||
}
|
||||
|
||||
func expectFrameLen(t *testing.T, frames <-chan int, want int) {
|
||||
t.Helper()
|
||||
select {
|
||||
case n := <-frames:
|
||||
if n != want {
|
||||
t.Fatalf("received frame len = %d, want %d", n, want)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("server did not receive frame in time")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebSocketRouteHandlerChecksBrowserOrigin 回归保护 Origin 修复:浏览器发起的 WS 升级
|
||||
// 必带 Origin(≠ Host),白名单来源需要改写 Origin 通过 coder/websocket,同名单外来源必须 403。
|
||||
func TestWebSocketRouteHandlerChecksBrowserOrigin(t *testing.T) {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
defer func() { _ = ln.Close() }()
|
||||
|
||||
_, wsHandler := transport.WebsocketListener(ln.Addr())
|
||||
httpServer := &http.Server{Handler: websocketRouteHandler(wsHandler, []string{"http://localhost:1234"})}
|
||||
go func() { _ = httpServer.Serve(ln) }()
|
||||
defer func() { _ = httpServer.Close() }()
|
||||
|
||||
host := ln.Addr().String()
|
||||
|
||||
// 白名单跨源升级:Origin 指向页面来源(端口/主机不同于 Host)。
|
||||
status := wsUpgradeStatus(t, host, "/apiws", "http://localhost:1234")
|
||||
if !strings.Contains(status, "101") {
|
||||
t.Fatalf("allowed cross-origin /apiws upgrade: got status %q, want 101 Switching Protocols", status)
|
||||
}
|
||||
|
||||
status = wsUpgradeStatus(t, host, "/apiws", "http://evil.example")
|
||||
if !strings.Contains(status, "403") {
|
||||
t.Fatalf("disallowed origin: got status %q, want 403", status)
|
||||
}
|
||||
|
||||
// 非白名单路径必须 404,不得升级。
|
||||
status = wsUpgradeStatus(t, host, "/nope", "http://localhost:1234")
|
||||
if !strings.Contains(status, "404") {
|
||||
t.Fatalf("disallowed path: got status %q, want 404", status)
|
||||
}
|
||||
}
|
||||
|
||||
// wsUpgradeStatus 用裸连接发一个合法的 WebSocket 升级请求并返回状态行。
|
||||
func wsUpgradeStatus(t *testing.T, host, path, origin string) string {
|
||||
t.Helper()
|
||||
conn, err := net.Dial("tcp", host)
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
defer func() { _ = conn.Close() }()
|
||||
|
||||
var keyBytes [16]byte
|
||||
if _, err := rand.Read(keyBytes[:]); err != nil {
|
||||
t.Fatalf("rand: %v", err)
|
||||
}
|
||||
key := base64.StdEncoding.EncodeToString(keyBytes[:])
|
||||
req := fmt.Sprintf("GET %s HTTP/1.1\r\nHost: %s\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n"+
|
||||
"Sec-WebSocket-Key: %s\r\nSec-WebSocket-Version: 13\r\nSec-WebSocket-Protocol: binary\r\nOrigin: %s\r\n\r\n",
|
||||
path, host, key, origin)
|
||||
|
||||
if err := conn.SetDeadline(time.Now().Add(5 * time.Second)); err != nil {
|
||||
t.Fatalf("set deadline: %v", err)
|
||||
}
|
||||
if _, err := conn.Write([]byte(req)); err != nil {
|
||||
t.Fatalf("write upgrade: %v", err)
|
||||
}
|
||||
statusLine, err := bufio.NewReader(conn).ReadString('\n')
|
||||
if err != nil {
|
||||
t.Fatalf("read status: %v", err)
|
||||
}
|
||||
return statusLine
|
||||
}
|
||||
|
||||
// TestServerObfuscatedTCPNotBlockedByStalledClient 回归保护「去 worker 池 / 握手移出 accept
|
||||
// 循环」的修复:一个发了几字节就挂起的连接,过去会卡死串行 accept 循环、阻塞所有后续接入;
|
||||
// 现在它只占用自己的 goroutine,正常客户端仍能在握手超时内被服务。
|
||||
func TestServerObfuscatedTCPNotBlockedByStalledClient(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, 1)
|
||||
srv := New(Options{Logger: zaptest.NewLogger(t), ObfuscatedTCP: true, WebSocket: true})
|
||||
srv.onFrame = func(n int) {
|
||||
select {
|
||||
case frames <- n:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
serveErr := make(chan error, 1)
|
||||
go func() { serveErr <- srv.Serve(ctx, ln) }()
|
||||
|
||||
// 挂起连接:发 4 个非 HTTP 字节通过分流(路由到 TCP),但不补满 obfuscated2 的 64 字节
|
||||
// init,使其卡在去混淆读取上。修复前这会拖死整个 TCP accept 循环。
|
||||
stalled, err := net.Dial("tcp", ln.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatalf("stalled dial: %v", err)
|
||||
}
|
||||
defer func() { _ = stalled.Close() }()
|
||||
if _, err := stalled.Write([]byte{0x01, 0x02, 0x03, 0x04}); err != nil {
|
||||
t.Fatalf("stalled write: %v", err)
|
||||
}
|
||||
|
||||
// 正常的 obfuscated abridged 客户端:应当照常被服务(onFrame 触发)。
|
||||
raw, err := net.Dial("tcp", ln.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatalf("tcp dial: %v", err)
|
||||
}
|
||||
obfs := obfuscator.Obfuscated2(rand.Reader, raw)
|
||||
if err := obfs.Handshake((codec.Abridged{}).ObfuscatedTag(), 2, mtproxy.Secret{}); err != nil {
|
||||
t.Fatalf("tcp obfuscated handshake: %v", err)
|
||||
}
|
||||
tcpConn, err := transport.NewProtocol(func() transport.Codec {
|
||||
return transport.Abridged.CodecNoHeader()
|
||||
}).Handshake(obfs)
|
||||
if err != nil {
|
||||
t.Fatalf("tcp transport handshake: %v", err)
|
||||
}
|
||||
|
||||
var payload bin.Buffer
|
||||
payload.PutInt32(0x12345678)
|
||||
payload.PutInt32(0x0badf00d)
|
||||
sendCtx, sc := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
if err := tcpConn.Send(sendCtx, &payload); err != nil {
|
||||
sc()
|
||||
t.Fatalf("tcp send: %v", err)
|
||||
}
|
||||
sc()
|
||||
expectFrameLen(t, frames, payload.Len())
|
||||
_ = tcpConn.Close()
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case err := <-serveErr:
|
||||
if err != nil {
|
||||
t.Fatalf("serve returned error: %v", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("server did not stop after ctx cancel")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue