Initial open source release
This commit is contained in:
commit
74992e893f
377 changed files with 118084 additions and 0 deletions
171
internal/mtprotoedge/exchange_test.go
Normal file
171
internal/mtprotoedge/exchange_test.go
Normal file
|
|
@ -0,0 +1,171 @@
|
|||
package mtprotoedge
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"go.uber.org/zap/zaptest"
|
||||
|
||||
"github.com/gotd/td/bin"
|
||||
"github.com/gotd/td/exchange"
|
||||
"github.com/gotd/td/mt"
|
||||
tgproto "github.com/gotd/td/proto"
|
||||
"github.com/gotd/td/transport"
|
||||
|
||||
"telesrv/internal/store"
|
||||
"telesrv/internal/store/memory"
|
||||
)
|
||||
|
||||
// TestKeyExchange 验证 M1:client 用 server 公钥完成 MTProto 密钥交换,
|
||||
// 双方得到一致的 auth key 与 server salt,且 server 将其存入 AuthKeyStore。
|
||||
func TestKeyExchange(t *testing.T) {
|
||||
const dc = 2
|
||||
|
||||
rsaKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
t.Fatalf("gen rsa: %v", err)
|
||||
}
|
||||
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
|
||||
keys := memory.NewAuthKeyStore()
|
||||
srv := New(Options{
|
||||
Logger: zaptest.NewLogger(t),
|
||||
DC: dc,
|
||||
RSAKey: rsaKey,
|
||||
AuthKeys: keys,
|
||||
})
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
serveErr := make(chan error, 1)
|
||||
go func() { serveErr <- srv.Serve(ctx, ln) }()
|
||||
|
||||
// client:TCP 拨号 + intermediate 握手,跑 client 端密钥交换。
|
||||
raw, err := net.Dial("tcp", ln.Addr().String())
|
||||
if err != nil {
|
||||
t.Fatalf("dial: %v", err)
|
||||
}
|
||||
conn, err := transport.Intermediate.Handshake(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("transport handshake: %v", err)
|
||||
}
|
||||
|
||||
pub := exchange.PublicKey{RSA: &rsaKey.PublicKey}
|
||||
exchCtx, ec := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer ec()
|
||||
res, err := exchange.NewExchanger(conn, dc).
|
||||
WithRand(rand.Reader).
|
||||
WithLogger(zaptest.NewLogger(t).Named("client")).
|
||||
Client([]exchange.PublicKey{pub}).
|
||||
Run(exchCtx)
|
||||
if err != nil {
|
||||
t.Fatalf("client exchange: %v", err)
|
||||
}
|
||||
|
||||
// server 在 Run 返回后落库,轮询等待。
|
||||
var saved store.AuthKeyData
|
||||
found := false
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for time.Now().Before(deadline) {
|
||||
saved, found, _ = keys.Get(context.Background(), res.AuthKey.ID)
|
||||
if found {
|
||||
break
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("server did not store auth key %x", res.AuthKey.ID)
|
||||
}
|
||||
if saved.Value != [256]byte(res.AuthKey.Value) {
|
||||
t.Fatal("server auth key value mismatch")
|
||||
}
|
||||
if saved.ServerSalt != res.ServerSalt {
|
||||
t.Fatalf("server salt mismatch: server=%d client=%d", saved.ServerSalt, res.ServerSalt)
|
||||
}
|
||||
|
||||
cancel()
|
||||
select {
|
||||
case err := <-serveErr:
|
||||
if err != nil {
|
||||
t.Fatalf("serve: %v", err)
|
||||
}
|
||||
case <-time.After(5 * time.Second):
|
||||
t.Fatal("server did not stop after ctx cancel")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReconnectFakeReqPQThenEncryptedFrame(t *testing.T) {
|
||||
const dc = 2
|
||||
addr, pub, _ := startTestServer(t, Options{DC: dc})
|
||||
|
||||
firstConn, auth, cipher := dialHandshake(t, addr, dc, pub)
|
||||
_ = firstConn.Close()
|
||||
|
||||
raw, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatalf("dial reconnect: %v", err)
|
||||
}
|
||||
conn, err := transport.Intermediate.Handshake(raw)
|
||||
if err != nil {
|
||||
t.Fatalf("transport reconnect: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
|
||||
var reqPayload bin.Buffer
|
||||
nonce, err := randInt128ForTest()
|
||||
if err != nil {
|
||||
t.Fatalf("nonce: %v", err)
|
||||
}
|
||||
if err := (&mt.ReqPqMultiRequest{Nonce: nonce}).Encode(&reqPayload); err != nil {
|
||||
t.Fatalf("encode req_pq_multi: %v", err)
|
||||
}
|
||||
var fakeReq bin.Buffer
|
||||
if err := (tgproto.UnencryptedMessage{
|
||||
MessageID: int64(tgproto.NewMessageID(time.Now(), tgproto.MessageFromClient)),
|
||||
MessageData: reqPayload.Raw(),
|
||||
}).Encode(&fakeReq); err != nil {
|
||||
t.Fatalf("encode fake req_pq: %v", err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
if err := conn.Send(ctx, &fakeReq); err != nil {
|
||||
cancel()
|
||||
t.Fatalf("send fake req_pq: %v", err)
|
||||
}
|
||||
cancel()
|
||||
|
||||
msgGen := tgproto.NewMessageIDGen(time.Now)
|
||||
sendEncrypted(t, conn, cipher, auth, msgGen.New(tgproto.MessageFromClient), &mt.PingRequest{PingID: 7})
|
||||
|
||||
var resPQFrame bin.Buffer
|
||||
ctx, cancel = context.WithTimeout(context.Background(), 5*time.Second)
|
||||
err = conn.Recv(ctx, &resPQFrame)
|
||||
cancel()
|
||||
if err != nil {
|
||||
t.Fatalf("recv resPQ: %v", err)
|
||||
}
|
||||
var plain tgproto.UnencryptedMessage
|
||||
if err := plain.Decode(&resPQFrame); err != nil {
|
||||
t.Fatalf("decode resPQ frame: %v", err)
|
||||
}
|
||||
if id, err := (&bin.Buffer{Buf: plain.MessageData}).PeekID(); err != nil || id != mt.ResPQTypeID {
|
||||
t.Fatalf("resPQ payload id = %#x err=%v, want %#x", id, err, mt.ResPQTypeID)
|
||||
}
|
||||
|
||||
got := collectReplies(t, conn, cipher, auth.AuthKey, mt.PongTypeID)
|
||||
mustHave(t, got, mt.PongTypeID, "pong after fake req_pq reconnect")
|
||||
}
|
||||
|
||||
func randInt128ForTest() (v bin.Int128, err error) {
|
||||
_, err = rand.Read(v[:])
|
||||
return v, err
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue