owpengram-server/internal/mtprotoedge/encrypted.go

1309 lines
44 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package mtprotoedge
import (
"bytes"
"compress/gzip"
"context"
"crypto/sha256"
"encoding/binary"
"errors"
"fmt"
"io"
"math"
"sync/atomic"
"time"
"go.uber.org/zap"
"github.com/gotd/td/bin"
"github.com/gotd/td/crypto"
"github.com/gotd/td/mt"
"github.com/gotd/td/proto"
"github.com/gotd/td/tgerr"
"github.com/gotd/td/transport"
"telesrv/internal/compat/layerwire"
"telesrv/internal/observability/dbtrace"
"telesrv/internal/postresponse"
"telesrv/internal/store"
)
// connState 是单连接的 MTProto 运行态。
type connState struct {
sentCreated bool
seen map[int64]clientMsgRecord // 已处理的 client msg_id用于幂等和 msgs_state_req
order []int64
minSeen int64
maxSeen int64
// maxContentMsgID/maxContentSeqNo 是已接受 content 消息的 msg_id / seq_no 高水位,
// 供 validateSeq 的 O(1) 快路径使用(客户端正常发送严格递增)。二者只增不减、
// 不随 seen 淘汰回退——快路径只接受「全扫描也必然接受」的子集,其余回落全扫描。
maxContentMsgID int64
maxContentSeqNo int32
}
type clientMsgRecord struct {
state byte
seqNo int32
content bool
}
func newConnState() *connState {
return &connState{
seen: make(map[int64]clientMsgRecord),
minSeen: math.MaxInt64,
}
}
func (cs *connState) reset() {
next := newConnState()
*cs = *next
}
const (
maxTrackedClientMsgIDs = 400
// maxContainerMessages bounds per-frame recursive work and ack growth. Official clients batch
// far fewer messages; 1024 leaves ample headroom while preventing a 16 MiB frame of zero-body
// container entries from expanding into tens of MiB of Go objects.
maxContainerMessages = 1024
// maxDispatchDepth bounds gzip/container wrapper recursion. Normal shapes are RPC, gzip(RPC),
// container(RPC...) and gzip(container(...)); deeper nesting has no compatibility value.
maxDispatchDepth = 4
// gotd already caps each gzip expansion at 10 MiB. This cumulative cap prevents several nested
// gzip layers in one transport frame from repeatedly allocating/decompressing that allowance.
maxDispatchExpandedBytes = 32 << 20
maxSingleGZIPExpandedBytes = 10 << 20
// MTProto service vectors operate on bounded connection tracking tables. Accepting more IDs
// only burns decode/CPU and cannot improve the result.
maxServiceMessageIDs = 4096
// A decoded container descriptor is 48 bytes on 64-bit Go today. Charge 64 bytes per entry
// before allocating the exact-size slice so allocator rounding and future field growth remain
// inside the process-wide inbound budget. Message bodies stay as zero-copy views of the already
// charged plaintext frame/gzip expansion.
containerDescriptorBudgetBytes = 64
msgStateUnknown byte = 1
msgStateNotReceived byte = 2
msgStateNotReceivedHigh byte = 3
msgStateReceived byte = 4
badMsgIDTooLow = 16
badMsgIDTooHigh = 17
badMsgIDInvalidBits = 18
badMsgSeqTooLow = 32
badMsgSeqTooHigh = 33
badMsgSeqNotEven = 34
badMsgSeqNotOdd = 35
badMsgContainer = 64
)
// handleEncrypted 解密加密消息,按需注册连接,处理服务消息并分发明文 payload。
// 返回(可能新建/更新的)当前连接对象,供 serveConn 维护生命周期。
// fetchedKey 非 nil 表示本帧的 auth key 是刚从 AuthKeyStore 查出的(首帧/换 auth key/被销毁
// 后回落);为 nil 表示走快路径——serveConn 判定 current 仍持同一未销毁的 auth key直接复用
// current.key/current.salt 解密,既不回查 AuthKeyStore 也不重建 store.AuthKeyData。
// plain 是 serveConn 持有的复用明文缓冲frame 的 slice 仅在下一帧解密前有效。
func (s *Server) handleEncrypted(ctx context.Context, tc transport.Conn, cs *connState, current *Conn, fetchedKey *store.AuthKeyData, b, plain *bin.Buffer) (*Conn, error) {
var key crypto.AuthKey
var serverSalt int64
if fetchedKey != nil {
key = crypto.AuthKey{Value: crypto.Key(fetchedKey.Value), ID: fetchedKey.ID}
serverSalt = fetchedKey.ServerSalt
} else {
// 快路径:复用已建立连接缓存的密钥与盐(同一 auth key 的后续帧,含同连接换 session
key = current.key
serverSalt = current.salt
}
frame, err := decryptClientFrame(key, b, plain)
if err != nil {
return current, fmt.Errorf("decrypt: %w", err)
}
if frame.salt != serverSalt {
c := current
temp := false
if c == nil || c.sessionID != frame.sessionID || c.authKeyID != key.ID {
c = s.newConn(tc, key, frame.sessionID, serverSalt)
temp = true
}
err := s.sendBadServerSalt(ctx, c, frame.messageID, frame.seqNo, serverSalt)
if temp {
c.Close()
}
return current, err
}
// 首个加密消息或 session 变化时(重新)注册连接到 SessionManager。
if current == nil || current.sessionID != frame.sessionID || current.authKeyID != key.ID {
if current != nil {
cs.reset()
}
if current != nil {
s.conns.Unregister(current)
current.Close()
}
current = s.newConn(tc, key, frame.sessionID, serverSalt)
// 注册即播种协商 layer新 Conn 的 clientLayer 为 0=canonical 227若等到
// 首条 RPC 的 Dispatch 返回后才刷新,重连老客户端在首条 RPC handler 执行期间
// 收到的 pending flush / 并发 push 会漏降级。进程内重连时 rpc 层留有
// (auth_key, session) / auth_key 两级协商记录,这里一次查询即可闭合该空窗。
if s.rpc != nil {
if layer, ok := s.rpc.NegotiatedLayer(current.authKeyID, current.sessionID); ok {
current.SetClientLayer(layer)
}
}
s.conns.Register(current)
}
s.maybePersistSession(ctx, current, frame.sessionID, key.ID, serverSalt)
body := frame.data
typeID, err := (&bin.Buffer{Buf: body}).PeekID()
if err != nil {
return current, fmt.Errorf("peek encrypted payload type id: %w", err)
}
if code := validateClientEnvelope(s.clock.Now(), frame.messageID, frame.seqNo, typeID); code != 0 {
s.log.Debug("Sending bad_msg_notification",
zap.Int64("msg_id", frame.messageID),
zap.Int32("seq_no", frame.seqNo),
zap.Uint32("type_id", typeID),
zap.Int("code", code),
)
return current, s.sendBadMsg(ctx, current, frame.messageID, frame.seqNo, code)
}
if err := sendQuickAckIfRequested(ctx, tc, key, frame.plaintext, s.writeTimeout); err != nil {
return current, err
}
content := clientMessageNeedsAck(typeID)
if record, ok := cs.seenRecord(frame.messageID); ok {
s.log.Debug("Duplicate msg_id; replay cached result if available", zap.Int64("msg_id", frame.messageID))
if err := s.replayRPCResultByRequest(ctx, current, frame.messageID); err != nil {
return current, err
}
if !record.content {
return current, nil
}
return current, s.sendAck(ctx, current, frame.messageID)
}
if code := cs.validateSeq(frame.messageID, frame.seqNo, content); code != 0 {
s.log.Debug("Sending bad_msg_notification",
zap.Int64("msg_id", frame.messageID),
zap.Int32("seq_no", frame.seqNo),
zap.Uint32("type_id", typeID),
zap.Int("code", code),
)
return current, s.sendBadMsg(ctx, current, frame.messageID, frame.seqNo, code)
}
cs.track(frame.messageID, frame.seqNo, content, msgStateReceived)
if !cs.sentCreated {
cs.sentCreated = true
s.log.Debug("Sending new_session_created", zap.Int64("msg_id", frame.messageID), zap.Int32("seq_no", frame.seqNo))
if err := s.sendNewSessionCreated(ctx, current, frame.messageID); err != nil {
return current, err
}
}
var acks []int64
if err := s.dispatch(ctx, cs, current, frame.messageID, frame.seqNo, &bin.Buffer{Buf: body}, &acks); err != nil {
return current, err
}
if len(acks) > 0 {
if err := s.sendAck(ctx, current, acks...); err != nil {
return current, err
}
}
return current, nil
}
// sessionSaveMinInterval 是单连接持久化 session 记录的最小间隔。把原本「每帧一次 Redis SET」
// 去抖到固定间隔——session 是软状态(生产无热读路径),只需周期刷新 last_seen/续 TTL。
const sessionSaveMinInterval = 30 * time.Second
// maybePersistSession 按 sessionSaveMinInterval 去抖持久化 session失败只告警不断连。
// 原实现每帧同步 Save 且失败即断连N 连接×帧率的 Redis 写放大 + Redis 抖动级联断连。
func (s *Server) maybePersistSession(ctx context.Context, c *Conn, sessionID int64, authKeyID [8]byte, salt int64) {
if c == nil {
return
}
now := s.clock.Now().Unix()
if last := c.lastSessionSaveUnix.Load(); last != 0 && now-last < int64(sessionSaveMinInterval/time.Second) {
return
}
c.lastSessionSaveUnix.Store(now)
if err := s.sessions.Save(ctx, store.SessionData{
ID: sessionID,
AuthKeyID: authKeyID,
Salt: salt,
LastSeen: now,
}); err != nil {
s.log.Warn("Persist session failed (non-fatal)",
zap.Int64("session_id", sessionID),
zap.Error(err),
)
}
}
func sendQuickAckIfRequested(ctx context.Context, tc transport.Conn, key crypto.AuthKey, plaintext []byte, writeTimeout time.Duration) error {
q, ok := tc.(quickAckTransport)
if !ok || !q.ConsumeQuickAckRequested() {
return nil
}
token := clientQuickAckToken(key, plaintext)
deadline := time.Time{}
if writeTimeout > 0 {
deadline = time.Now().Add(writeTimeout)
}
if d, ok := ctx.Deadline(); ok && (deadline.IsZero() || d.Before(deadline)) {
deadline = d
}
if dq, ok := tc.(deadlineQuickAckTransport); ok {
return dq.SendQuickAckDeadline(deadline, token)
}
if deadline.IsZero() {
return q.SendQuickAck(ctx, token)
}
sendCtx, cancel := context.WithDeadline(ctx, deadline)
defer cancel()
return q.SendQuickAck(sendCtx, token)
}
// clientQuickAckToken 按 Android MTProto v2 公式计算 quick ackSHA256(auth_key[88:120] +
// 完整明文)[:4]。plaintext 直接来自解密复用缓冲decryptClientFrame.plaintext
// 与旧实现「把解密结果重编码一遍再哈希」字节一致但零拷贝。
func clientQuickAckToken(key crypto.AuthKey, plaintext []byte) uint32 {
h := sha256.New()
_, _ = h.Write(key.Value[88:120])
_, _ = h.Write(plaintext)
sum := h.Sum(nil)
return binary.LittleEndian.Uint32(sum[:4]) &^ quickAckResponseFlag
}
// dispatch 处理一条明文消息:解包 container/gzip处理服务消息其余转 RPC 路由。
// content-related 消息ping、RPC的 msg_id 会收集到 acks 以便统一确认。
func (s *Server) dispatch(ctx context.Context, cs *connState, c *Conn, msgID int64, seqNo int32, b *bin.Buffer, acks *[]int64) error {
expanded := 0
return s.dispatchWithBudget(ctx, cs, c, msgID, seqNo, b, acks, dispatchBudget{expanded: &expanded})
}
type dispatchBudget struct {
depth int
containerDepth int
expanded *int
}
func (s *Server) dispatchWithBudget(ctx context.Context, cs *connState, c *Conn, msgID int64, seqNo int32, b *bin.Buffer, acks *[]int64, budget dispatchBudget) error {
if budget.depth > maxDispatchDepth {
return fmt.Errorf("mtproto wrapper depth %d exceeds %d", budget.depth, maxDispatchDepth)
}
id, err := b.PeekID()
if err != nil {
return fmt.Errorf("peek type id: %w", err)
}
ackContent := func() {
if clientMessageNeedsAck(id) {
*acks = append(*acks, msgID)
}
}
switch id {
case proto.GZIPTypeID:
data, releaseExpansion, err := s.decodeGZIPWithGlobalBudget(b)
if err != nil {
return fmt.Errorf("decode gzip: %w", err)
}
defer releaseExpansion()
*budget.expanded += len(data)
if *budget.expanded > maxDispatchExpandedBytes {
return fmt.Errorf("cumulative gzip expansion %d exceeds %d", *budget.expanded, maxDispatchExpandedBytes)
}
budget.depth++
return s.dispatchWithBudget(ctx, cs, c, msgID, seqNo, &bin.Buffer{Buf: data}, acks, budget)
case proto.MessageContainerTypeID:
if budget.containerDepth != 0 {
return s.sendBadMsg(ctx, c, msgID, seqNo, badMsgContainer)
}
count, err := containerMessageCount(b)
if err != nil {
return fmt.Errorf("decode container count: %w", err)
}
if count > maxContainerMessages {
return s.sendBadMsg(ctx, c, msgID, seqNo, badMsgContainer)
}
container, releaseContainer, err := s.decodeMessageContainerViews(b, count)
if err != nil {
return fmt.Errorf("decode container: %w", err)
}
defer releaseContainer()
if code := validateClientContainer(msgID, seqNo, container); code != 0 {
return s.sendBadMsg(ctx, c, msgID, seqNo, code)
}
budget.depth++
budget.containerDepth++
for i := range container.Messages {
m := container.Messages[i]
typeID, err := (&bin.Buffer{Buf: m.Body}).PeekID()
if err != nil {
return fmt.Errorf("peek container message type id: %w", err)
}
content := clientMessageNeedsAck(typeID)
if record, ok := cs.seenRecord(m.ID); ok {
if err := s.replayRPCResultByRequest(ctx, c, m.ID); err != nil {
return err
}
if record.content {
*acks = append(*acks, m.ID)
}
continue
}
if code := cs.validateSeq(m.ID, int32(m.SeqNo), content); code != 0 {
return s.sendBadMsg(ctx, c, m.ID, int32(m.SeqNo), code)
}
cs.track(m.ID, int32(m.SeqNo), content, msgStateReceived)
if err := s.dispatchWithBudget(ctx, cs, c, m.ID, int32(m.SeqNo), &bin.Buffer{Buf: m.Body}, acks, budget); err != nil {
return err
}
}
return nil
case mt.PingRequestTypeID:
var ping mt.PingRequest
if err := ping.Decode(b); err != nil {
return fmt.Errorf("decode ping: %w", err)
}
ackContent()
return s.sendPong(ctx, c, msgID, ping.PingID)
case mt.PingDelayDisconnectRequestTypeID:
var ping mt.PingDelayDisconnectRequest
if err := ping.Decode(b); err != nil {
return fmt.Errorf("decode ping_delay_disconnect: %w", err)
}
ackContent()
return s.sendPong(ctx, c, msgID, ping.PingID)
case mt.GetFutureSaltsRequestTypeID:
var req mt.GetFutureSaltsRequest
if err := req.Decode(b); err != nil {
return fmt.Errorf("decode get_future_salts: %w", err)
}
ackContent()
return s.sendFutureSalts(ctx, c, msgID, req.Num)
case mt.MsgsAckTypeID:
if err := validateFirstVectorCount(b, maxServiceMessageIDs); err != nil {
return fmt.Errorf("msgs_ack vector: %w", err)
}
var ack mt.MsgsAck
if err := ack.Decode(b); err != nil {
return fmt.Errorf("decode msgs_ack: %w", err)
}
c.AckServerMessages(ack.MsgIDs)
s.log.Debug("Received msgs_ack", zap.Int64s("msg_ids", ack.MsgIDs))
return nil
case mt.MsgsStateReqTypeID:
if err := validateFirstVectorCount(b, maxServiceMessageIDs); err != nil {
return fmt.Errorf("msgs_state_req vector: %w", err)
}
var req mt.MsgsStateReq
if err := req.Decode(b); err != nil {
return fmt.Errorf("decode msgs_state_req: %w", err)
}
ackContent()
outgoing, err := c.OutgoingStateInfo(ctx, req.MsgIDs)
if err != nil {
return err
}
return s.sendMsgsStateInfo(ctx, c, msgID, mergeStateInfo(outgoing, cs.stateInfo(req.MsgIDs)))
case mt.MsgResendReqTypeID:
if err := validateFirstVectorCount(b, maxServiceMessageIDs); err != nil {
return fmt.Errorf("msg_resend_req vector: %w", err)
}
var req mt.MsgResendReq
if err := req.Decode(b); err != nil {
return fmt.Errorf("decode msg_resend_req: %w", err)
}
ackContent()
outgoing, err := c.ResendMessages(ctx, req.MsgIDs)
if err != nil {
return err
}
return s.sendMsgsStateInfo(ctx, c, msgID, mergeStateInfo(outgoing, cs.stateInfo(req.MsgIDs)))
case mt.MsgsStateInfoTypeID:
reqMsgID, info, err := msgsStateInfoView(b)
if err != nil {
return fmt.Errorf("decode msgs_state_info: %w", err)
}
s.log.Debug("Received msgs_state_info", zap.Int64("req_msg_id", reqMsgID), zap.Int("len", len(info)))
return nil
case mt.MsgsAllInfoTypeID:
count, info, err := msgsAllInfoView(b)
if err != nil {
return fmt.Errorf("decode msgs_all_info: %w", err)
}
if len(info) != count {
return fmt.Errorf("decode msgs_all_info: info length %d does not match msg_ids %d", len(info), count)
}
s.log.Debug("Received msgs_all_info", zap.Int("msg_ids", count), zap.Int("len", len(info)))
return nil
case mt.DestroySessionRequestTypeID:
var req mt.DestroySessionRequest
if err := req.Decode(b); err != nil {
return fmt.Errorf("decode destroy_session: %w", err)
}
ackContent()
return s.sendDestroySession(ctx, c, req.SessionID)
case mt.HTTPWaitRequestTypeID:
var req mt.HTTPWaitRequest
if err := req.Decode(b); err != nil {
return fmt.Errorf("decode http_wait: %w", err)
}
s.log.Debug("Received http_wait",
zap.Int("max_delay", req.MaxDelay),
zap.Int("wait_after", req.WaitAfter),
zap.Int("max_wait", req.MaxWait),
)
return nil
case mt.RPCDropAnswerRequestTypeID:
var req mt.RPCDropAnswerRequest
if err := req.Decode(b); err != nil {
return fmt.Errorf("decode rpc_drop_answer: %w", err)
}
ackContent()
s.log.Debug("Received rpc_drop_answer", zap.Int64("req_msg_id", req.ReqMsgID))
return s.sendResult(ctx, c, msgID, &mt.RPCAnswerUnknown{})
case destroyAuthKeyRequestTypeID:
var req destroyAuthKeyRequest
if err := req.Decode(b); err != nil {
return err
}
ackContent()
s.log.Debug("Received destroy_auth_key", zap.String("auth_key_id", c.authKeyHex))
// 真正销毁:删密钥库记录(每帧回查,删除后该 key 的入站帧立即失效)并主动
// 断开同 key 的其他连接——出站推送用连接持有的密钥副本加密、不回查密钥库,
// 不断开的话被销毁 key 的空闲连接仍能持续收到推送。发起连接除外:响应要
// 先送达它的下一帧会因密钥缺失自然断开。授权authorizations不在此清理
// destroy_auth_key 是 PFS 密钥轮换的清理动作,不等于登出。
if err := s.authKeys.Delete(ctx, c.authKeyID); err != nil {
s.log.Warn("Delete auth key failed", zap.String("auth_key_id", c.authKeyHex), zap.Error(err))
return c.SendAsync(ctx, proto.MessageServerResponse, &destroyAuthKeyFail{})
}
// 标记密钥已销毁:发起连接被 CloseSessionsForRawAuthKeyExcept 排除(响应需先送达),
// 它下一帧不能再走 serveConn 的密钥复用快路径,须回落到 Get→AuthKeyNotFound 自然失效。
c.keyDestroyed.Store(true)
s.conns.CloseSessionsForRawAuthKeyExcept(c.authKeyID, c.sessionID)
return c.SendAsync(ctx, proto.MessageServerResponse, &destroyAuthKeyOk{})
default:
ackContent()
return s.enqueueRPC(ctx, c, msgID, id, b)
}
}
// decodeGZIPWithGlobalBudget reserves the maximum single-wrapper output before
// decompression starts. Once the actual size is known the excess reservation is
// returned, while the actual output remains charged through recursive dispatch.
// This closes the gap where every connection read goroutine could otherwise hold
// an unaccounted 10 MiB expansion before the shared RPC scheduler saw the body.
func (s *Server) decodeGZIPWithGlobalBudget(b *bin.Buffer) ([]byte, func(), error) {
compressed, err := gzipPackedBytesView(b)
if err != nil {
return nil, func() {}, err
}
reserved := int64(0)
release := func() {
if reserved > 0 && s.frameBudget != nil {
s.frameBudget.release(reserved)
reserved = 0
}
}
if s.frameBudget != nil {
reserved, err = s.frameBudget.reserve(maxSingleGZIPExpandedBytes, 0)
if err != nil {
return nil, func() {}, err
}
}
r, err := gzip.NewReader(bytes.NewReader(compressed))
if err != nil {
release()
return nil, func() {}, err
}
data, readErr := io.ReadAll(io.LimitReader(r, maxSingleGZIPExpandedBytes+1))
closeErr := r.Close()
if readErr != nil {
release()
return nil, func() {}, readErr
}
if closeErr != nil {
release()
return nil, func() {}, closeErr
}
if len(data) > maxSingleGZIPExpandedBytes {
release()
return nil, func() {}, fmt.Errorf("gzip expansion %d exceeds %d", len(data), maxSingleGZIPExpandedBytes)
}
if reserved > int64(len(data)) {
s.frameBudget.release(reserved - int64(len(data)))
reserved = int64(len(data))
}
return data, release, nil
}
// gzipPackedBytesView parses the TL bytes envelope without copying the compressed
// payload. proto.GZIP.Decode calls bin.Buffer.Bytes, which duplicates the compressed
// frame before allocating the decompressed result.
func gzipPackedBytesView(b *bin.Buffer) ([]byte, error) {
if b == nil || len(b.Buf) < 5 {
return nil, io.ErrUnexpectedEOF
}
if binary.LittleEndian.Uint32(b.Buf[:4]) != proto.GZIPTypeID {
return nil, fmt.Errorf("unexpected gzip constructor %#x", binary.LittleEndian.Uint32(b.Buf[:4]))
}
payload, _, err := tlBytesView(b.Buf[4:], -1)
return payload, err
}
// tlBytesView validates one TL bytes envelope and returns a view into the caller-owned buffer.
// maxPayload < 0 means that the enclosing frame budget is the only size limit. The limit is
// checked from the encoded length before touching the payload, so service messages cannot make
// generated decoders allocate an attacker-selected []byte first and validate it afterwards.
func tlBytesView(raw []byte, maxPayload int) ([]byte, int, error) {
if len(raw) < 1 {
return nil, 0, io.ErrUnexpectedEOF
}
header, size := 1, int(raw[0])
if size == 254 {
if len(raw) < 4 {
return nil, 0, io.ErrUnexpectedEOF
}
header = 4
size = int(raw[1]) | int(raw[2])<<8 | int(raw[3])<<16
} else if size == 255 {
return nil, 0, errors.New("invalid TL bytes length marker 255")
}
if maxPayload >= 0 && size > maxPayload {
return nil, 0, fmt.Errorf("TL bytes length %d exceeds %d", size, maxPayload)
}
padded := (header + size + 3) &^ 3
if size < 0 || padded < header || len(raw) < padded {
return nil, 0, io.ErrUnexpectedEOF
}
return raw[header : header+size : header+size], padded, nil
}
// decodeMessageContainerViews parses the container without proto.Message.Decode's per-body
// copies. Bodies are immutable views of b and stay alive only for this synchronous dispatch;
// enqueueRPC takes its own budgeted copy before returning. Only the exact-size descriptor slice
// is new memory, and that allocation is reserved globally first.
func (s *Server) decodeMessageContainerViews(b *bin.Buffer, count int) (proto.MessageContainer, func(), error) {
release := func() {}
if b == nil || len(b.Buf) < 8 {
return proto.MessageContainer{}, release, io.ErrUnexpectedEOF
}
if got := binary.LittleEndian.Uint32(b.Buf[:4]); got != proto.MessageContainerTypeID {
return proto.MessageContainer{}, release, fmt.Errorf("unexpected constructor %#x", got)
}
declared := int(int32(binary.LittleEndian.Uint32(b.Buf[4:8])))
if declared != count || count < 0 || count > maxContainerMessages {
return proto.MessageContainer{}, release, fmt.Errorf("invalid message count %d", declared)
}
reserved := int64(0)
if count > 0 && s.frameBudget != nil {
var err error
reserved, err = s.frameBudget.reserve(int64(count*containerDescriptorBudgetBytes), 0)
if err != nil {
return proto.MessageContainer{}, release, err
}
release = func() {
if reserved > 0 {
s.frameBudget.release(reserved)
reserved = 0
}
}
}
messages := make([]proto.Message, count)
offset := 8
for i := range messages {
if len(b.Buf)-offset < 16 {
release()
return proto.MessageContainer{}, func() {}, io.ErrUnexpectedEOF
}
id := int64(binary.LittleEndian.Uint64(b.Buf[offset : offset+8]))
seqNo := int32(binary.LittleEndian.Uint32(b.Buf[offset+8 : offset+12]))
bodyLen := int(int32(binary.LittleEndian.Uint32(b.Buf[offset+12 : offset+16])))
offset += 16
if bodyLen < 0 || bodyLen > 1024*1024 {
release()
return proto.MessageContainer{}, func() {}, fmt.Errorf("message length %d is invalid", bodyLen)
}
if bodyLen > len(b.Buf)-offset {
release()
return proto.MessageContainer{}, func() {}, io.ErrUnexpectedEOF
}
bodyEnd := offset + bodyLen
messages[i] = proto.Message{
ID: id,
SeqNo: int(seqNo),
Bytes: bodyLen,
Body: b.Buf[offset:bodyEnd:bodyEnd],
}
offset = bodyEnd
}
return proto.MessageContainer{Messages: messages}, release, nil
}
func msgsStateInfoView(b *bin.Buffer) (int64, []byte, error) {
if b == nil || len(b.Buf) < 12 {
return 0, nil, io.ErrUnexpectedEOF
}
if got := binary.LittleEndian.Uint32(b.Buf[:4]); got != mt.MsgsStateInfoTypeID {
return 0, nil, fmt.Errorf("unexpected constructor %#x", got)
}
info, _, err := tlBytesView(b.Buf[12:], maxServiceMessageIDs)
if err != nil {
return 0, nil, err
}
return int64(binary.LittleEndian.Uint64(b.Buf[4:12])), info, nil
}
func msgsAllInfoView(b *bin.Buffer) (int, []byte, error) {
if err := validateFirstVectorCount(b, maxServiceMessageIDs); err != nil {
return 0, nil, fmt.Errorf("vector: %w", err)
}
count := int(int32(binary.LittleEndian.Uint32(b.Buf[8:12])))
// count is already non-negative and capped, but check remaining bytes before multiplying into
// an offset so malformed frames cannot produce an out-of-bounds slice.
if count > (len(b.Buf)-12)/8 {
return 0, nil, io.ErrUnexpectedEOF
}
offset := 12 + count*8
info, _, err := tlBytesView(b.Buf[offset:], maxServiceMessageIDs)
if err != nil {
return 0, nil, err
}
return count, info, nil
}
func containerMessageCount(b *bin.Buffer) (int, error) {
if b == nil || len(b.Buf) < 8 {
return 0, io.ErrUnexpectedEOF
}
if binary.LittleEndian.Uint32(b.Buf[:4]) != proto.MessageContainerTypeID {
return 0, fmt.Errorf("unexpected constructor %#x", binary.LittleEndian.Uint32(b.Buf[:4]))
}
count := int(int32(binary.LittleEndian.Uint32(b.Buf[4:8])))
if count < 0 {
return 0, fmt.Errorf("negative message count %d", count)
}
return count, nil
}
func validateFirstVectorCount(b *bin.Buffer, max int) error {
if b == nil || len(b.Buf) < 12 {
return io.ErrUnexpectedEOF
}
if got := binary.LittleEndian.Uint32(b.Buf[4:8]); got != bin.TypeVector {
return fmt.Errorf("unexpected vector constructor %#x", got)
}
count := int(int32(binary.LittleEndian.Uint32(b.Buf[8:12])))
if count < 0 {
return fmt.Errorf("negative vector count %d", count)
}
if count > max {
return fmt.Errorf("vector count %d exceeds %d", count, max)
}
return nil
}
func mergeStateInfo(primary, fallback []byte) []byte {
if len(primary) == 0 {
return fallback
}
info := make([]byte, len(fallback))
copy(info, fallback)
for i, state := range primary {
if i >= len(info) {
break
}
if state != 0 {
info[i] = state
}
}
return info
}
// enqueueRPC 把一条 RPC 请求交给连接的 inbound 调度器。typeID 由 dispatch 传入
// (已 PeekID 过一次method 只解析一次并随任务透传,避免同一请求三处重复 PeekID/typeName。
func (s *Server) enqueueRPC(ctx context.Context, c *Conn, msgID int64, typeID uint32, request *bin.Buffer) error {
method := s.typeName(typeID)
if cached, ok := s.cachedRPCResult(c, msgID); ok {
s.log.Info("RPC duplicate replay from session cache",
zap.String("method", method),
zap.Int64("msg_id", msgID),
zap.String("auth_key_id", c.authKeyHex),
zap.Int64("session_id", c.sessionID),
)
return c.SendEncoded(ctx, proto.MessageServerResponse, cached)
}
// 两级条数/字节预算必须先于 Copy对抗客户端不能用大量满尺寸请求在“判断队列满”
// 之前制造一轮无上限的临时 body 分配。reservation 在 commit/abort 间唯一持有预算。
reservation, err := c.reserveInboundRPC(ctx, method, request.Len())
if err != nil {
return s.handleInboundRPCAdmissionError(ctx, c, msgID, method, err)
}
defer reservation.abort()
body := request.Copy()
responseGate := &rpcResponseGate{}
timeoutResponse := func() {
if !responseGate.tryTimeout() {
return
}
// 原 task context 已到期,使用有界的新 context 回显明确的可重试超时;
// 500 保持 TDesktop 默认重试语义,错误名区分于容量型 FLOOD_WAIT。
writeTimeout := c.writeTimeout
if writeTimeout <= 0 || writeTimeout > 5*time.Second {
writeTimeout = 5 * time.Second
}
responseCtx, cancel := context.WithTimeout(context.Background(), writeTimeout)
defer cancel()
if sendErr := s.sendResult(responseCtx, c, msgID, &mt.RPCError{
ErrorCode: 500,
ErrorMessage: "RPC_TIMEOUT",
}); sendErr != nil && !isClientDisconnect(sendErr) {
s.log.Debug("Send RPC timeout failed",
zap.String("method", method),
zap.Int64("msg_id", msgID),
zap.String("auth_key_id", c.authKeyHex),
zap.Int64("session_id", c.sessionID),
zap.Error(sendErr),
)
}
}
err = reservation.commit(inboundRPC{
method: method,
size: len(body),
onTimeout: timeoutResponse,
run: func(taskCtx context.Context) error {
// body 是预算成功后生成的独立副本,且每个任务只 run 一次,
// 无需再 append 拷贝;直接复用,省掉一份 inbound 在途内存。
if err := s.handleRPC(taskCtx, c, msgID, method, &bin.Buffer{Buf: body}, responseGate); err != nil {
fields := []zap.Field{
zap.Int64("msg_id", msgID),
zap.String("auth_key_id", c.authKeyHex),
zap.Int64("session_id", c.sessionID),
zap.Error(err),
}
if isClientDisconnect(err) {
s.log.Debug("RPC async handler canceled", fields...)
} else {
s.log.Info("RPC async handler failed", fields...)
}
return err
}
return nil
},
})
return s.handleInboundRPCAdmissionError(ctx, c, msgID, method, err)
}
func (s *Server) handleInboundRPCAdmissionError(ctx context.Context, c *Conn, msgID int64, method string, err error) error {
if errors.Is(err, ErrInboundRPCQueueFull) {
s.log.Debug("Inbound RPC capacity exhausted",
zap.String("method", method),
zap.Int64("msg_id", msgID),
zap.String("auth_key_id", c.authKeyHex),
zap.Int64("session_id", c.sessionID),
)
return s.sendResult(ctx, c, msgID, &mt.RPCError{
ErrorCode: 420,
ErrorMessage: "FLOOD_WAIT_1",
})
}
return err
}
// handleRPC 把明文 RPC 请求交给 RPC 路由,并将结果或错误包成 rpc_result 回发。
func (s *Server) handleRPC(ctx context.Context, c *Conn, msgID int64, method string, b *bin.Buffer, responseGate *rpcResponseGate) error {
if s.rpc == nil {
s.log.Warn("No RPC handler configured; dropping request", zap.String("method", method))
return nil
}
ctx = postresponse.WithCallbacks(ctx)
ctx, dbStats := dbtrace.WithStats(ctx)
start := s.clock.Now()
result, err := s.rpc.Dispatch(ctx, c.authKeyID, c.sessionID, b)
dur := s.clock.Now().Sub(start)
s.metrics.RPCHandled(method, dur, err)
// 刷新本连接协商 layerinvokeWithLayer/initConnection 已被 Dispatch 处理并登记),
// 供 rpc_result 与后续 push 出站降级使用。仅在确实观测到 layer 时更新——缓存被驱逐
// 时 NegotiatedLayer 返回 ok=false此时必须保留连接已记住的 layer绝不覆盖成默认值
// 否则长连接老客户端的条目被驱逐后会被误降回 227。
if layer, ok := s.rpc.NegotiatedLayer(c.authKeyID, c.sessionID); ok {
c.SetClientLayer(layer)
}
fields := make([]zap.Field, 0, 12)
fields = append(fields,
zap.String("method", method),
zap.String("auth_key_id", c.authKeyHex),
zap.Int64("session_id", c.sessionID),
zap.Int64("msg_id", msgID),
zap.Duration("dur", dur),
)
if businessAuthKeyHex, ok := c.BusinessAuthKeyHex(); ok {
fields = append(fields, zap.String("business_auth_key_id", businessAuthKeyHex))
}
if userID := c.UserID(); userID != 0 {
fields = append(fields, zap.Int64("user_id", userID))
}
fields = dbtrace.AppendZapFields(fields, "", dbStats.Snapshot())
if ctxErr := ctx.Err(); ctxErr != nil {
// A canceled request context means neither a success nor an error can be delivered
// with this expired context. In particular, do not cache a late successful result and
// hand it to outbound: a past write deadline would correctly poison that transport and
// could prevent the scheduler's fresh-context RPC_TIMEOUT response from being sent.
cancelFields := append(fields, zap.NamedError("context_error", ctxErr))
if err != nil {
cancelFields = append(cancelFields, zap.NamedError("dispatch_error", err))
}
s.log.Info("RPC canceled", cancelFields...)
return ctxErr
}
// A deadline callback may have already emitted RPC_TIMEOUT while Dispatch was returning.
// Claim the single normal-response slot before serializing any success/error rpc_result.
if responseGate != nil && !responseGate.tryNormal() {
s.log.Info("RPC result suppressed after timeout", fields...)
return context.DeadlineExceeded
}
if err != nil {
var rpcErr *tgerr.Error
if errors.As(err, &rpcErr) {
s.log.Info("RPC error", append(fields, zap.Int("code", rpcErr.Code), zap.String("error", rpcErr.Message))...)
return s.sendResult(ctx, c, msgID, &mt.RPCError{
ErrorCode: rpcErr.Code,
ErrorMessage: rpcErr.Message,
})
}
s.log.Info("RPC internal error", append(fields, zap.Error(err))...)
return s.sendResult(ctx, c, msgID, &mt.RPCError{
ErrorCode: 500,
ErrorMessage: "INTERNAL",
})
}
s.log.Info("RPC handled", fields...)
if err := s.sendResult(ctx, c, msgID, result); err != nil {
return err
}
postresponse.Run(ctx)
return nil
}
// rpcResponseGate guarantees exactly one terminal rpc_result per request. A running deadline
// races legitimately with a handler completing at the boundary; whichever path claims state
// first owns the response, and the other path becomes a no-op.
type rpcResponseGate struct {
state atomic.Uint32
}
func (g *rpcResponseGate) tryNormal() bool {
return g == nil || g.state.CompareAndSwap(0, 1)
}
func (g *rpcResponseGate) tryTimeout() bool {
return g != nil && g.state.CompareAndSwap(0, 2)
}
// sendResult 把 RPC 结果包成 rpc_result 并加密回发。
func (s *Server) sendResult(ctx context.Context, c *Conn, reqMsgID int64, result bin.Encoder) error {
encoded, err := s.encodeRPCResult(c, reqMsgID, result)
if err != nil {
return err
}
s.storeRPCResult(c, reqMsgID, encoded)
return c.SendEncoded(ctx, proto.MessageServerResponse, encoded)
}
// encodeRPCResult 编码 rpc_result。内层对象与 rpc_result 头type_id + req_msg_id
// 一次性编码进同一 buffer——旧实现先编码内层、再经 proto.Result.Encode 整体拷贝一遍,
// 每条响应多一份全量 body 拷贝。内层按连接协商 layer 降级layer==227 直通,零开销),
// 降级改写字节时才重建整条消息。降级失败 fail-safe记日志并发送 canonical 字节——
// 宁可老客户端对个别长尾对象渲染异常,也不让连接/流崩。
func (s *Server) encodeRPCResult(c *Conn, reqMsgID int64, result bin.Encoder) (*encodedOutboundMessage, error) {
const headerLen = 4 + 8 // rpc_result#f35c6d01 type_id + req_msg_id
var buf bin.Buffer
buf.PutID(proto.ResultTypeID)
buf.PutLong(reqMsgID)
if err := result.Encode(&buf); err != nil {
return nil, fmt.Errorf("encode rpc result: %w", err)
}
if layer := c.ClientLayer(); layer < layerwire.CanonicalLayer {
inner := buf.Buf[headerLen:]
if down, err := layerwire.Transcode(inner, layer); err != nil {
s.log.Warn("layerwire downgrade failed; sending canonical rpc_result",
zap.Int("layer", layer), zap.Int64("req_msg_id", reqMsgID), zap.Error(err))
} else if !sameBacking(down, inner) {
var rebuilt bin.Buffer
rebuilt.PutID(proto.ResultTypeID)
rebuilt.PutLong(reqMsgID)
rebuilt.Put(down)
buf = rebuilt
}
}
return &encodedOutboundMessage{
typeID: proto.ResultTypeID,
body: buf.Raw(),
reqMsgID: reqMsgID,
}, nil
}
func (s *Server) cachedRPCResult(c *Conn, reqMsgID int64) (*encodedOutboundMessage, bool) {
if s == nil || s.rpcResults == nil || c == nil {
return nil, false
}
return s.rpcResults.Get(c.authKeyID, c.sessionID, reqMsgID)
}
func (s *Server) replayRPCResultByRequest(ctx context.Context, c *Conn, reqMsgID int64) error {
if c == nil {
return nil
}
if resent, err := c.ResendByRequest(ctx, reqMsgID); err != nil {
return err
} else if resent {
s.log.Debug("Resent connection cached rpc_result for duplicate msg_id", zap.Int64("msg_id", reqMsgID))
return nil
}
if cached, ok := s.cachedRPCResult(c, reqMsgID); ok {
if err := c.SendEncoded(ctx, proto.MessageServerResponse, cached); err != nil {
return err
}
s.log.Debug("Resent session cached rpc_result for duplicate msg_id", zap.Int64("msg_id", reqMsgID))
}
return nil
}
func (s *Server) storeRPCResult(c *Conn, reqMsgID int64, encoded *encodedOutboundMessage) {
if s == nil || s.rpcResults == nil || c == nil {
return
}
s.rpcResults.Put(c.authKeyID, c.sessionID, reqMsgID, encoded)
}
// sendPong 回复 mt.PingRequest / mt.PingDelayDisconnectRequest。
func (s *Server) sendPong(ctx context.Context, c *Conn, reqMsgID, pingID int64) error {
return c.SendAsync(ctx, proto.MessageServerResponse, &mt.Pong{MsgID: reqMsgID, PingID: pingID})
}
// sendFutureSalts 回复 MTProto get_future_salts。
//
// 第一阶段只维护当前 auth key 的权威 server_salt因此返回当前 salt 的有效窗口。
// 后续如引入 salt rotation可在这里扩展为多条未来 salt。
func (s *Server) sendFutureSalts(ctx context.Context, c *Conn, reqMsgID int64, num int) error {
if num < 0 {
num = 0
}
if num > 1 {
num = 1
}
now := int(s.clock.Now().Unix())
salts := make([]mt.FutureSalt, 0, num)
if num == 1 {
salts = append(salts, mt.FutureSalt{
ValidSince: now - 300,
ValidUntil: now + 24*60*60,
Salt: c.salt,
})
}
return c.SendAsync(ctx, proto.MessageServerResponse, &mt.FutureSalts{
ReqMsgID: reqMsgID,
Now: now,
Salts: salts,
})
}
// sendNewSessionCreated 在连接首个加密消息后通知客户端新 session 已建立。
// unique_id 必须每个 server session 实例独立:客户端按 unique_id 去重,
// 复用同一值会让断线重连后的 new_session_created 被吞掉,错过的差分补拉
// Android 收到后才调 getDifference随之丢失。
func (s *Server) sendNewSessionCreated(ctx context.Context, c *Conn, firstMsgID int64) error {
return c.SendAsync(ctx, proto.MessageFromServer, &mt.NewSessionCreated{
FirstMsgID: firstMsgID,
UniqueID: s.newServerSessionUID(),
ServerSalt: c.salt,
})
}
func (s *Server) newServerSessionUID() int64 {
var b [8]byte
if _, err := io.ReadFull(s.rand, b[:]); err == nil {
return int64(binary.LittleEndian.Uint64(b[:]))
}
return s.clock.Now().UnixNano()
}
// sendAck 确认收到客户端 content-related 消息。
func (s *Server) sendAck(ctx context.Context, c *Conn, ids ...int64) error {
return c.SendAsync(ctx, proto.MessageFromServer, &mt.MsgsAck{MsgIDs: ids})
}
// sendMsgsStateInfo 回复 msgs_state_req/msg_resend_req。
func (s *Server) sendMsgsStateInfo(ctx context.Context, c *Conn, reqMsgID int64, info []byte) error {
return c.SendAsync(ctx, proto.MessageServerResponse, &mt.MsgsStateInfo{ReqMsgID: reqMsgID, Info: info})
}
func (s *Server) sendDestroySession(ctx context.Context, c *Conn, sessionID int64) error {
removed := false
if sessionID != c.sessionID {
removed = s.conns.DestroySessionForAuthKey(c.authKeyID, sessionID)
if err := s.sessions.Delete(ctx, sessionID); err != nil {
s.log.Debug("Delete session record failed",
zap.String("auth_key_id", c.authKeyHex),
zap.Int64("session_id", sessionID),
zap.Error(err),
)
}
}
if removed {
return c.Send(ctx, proto.MessageServerResponse, &mt.DestroySessionOk{SessionID: sessionID})
}
return c.Send(ctx, proto.MessageServerResponse, &mt.DestroySessionNone{SessionID: sessionID})
}
// sendBadMsg 通知客户端消息存在协议层错误msg_id/seqno 非法)。
func (s *Server) sendBadMsg(ctx context.Context, c *Conn, badMsgID int64, badSeqno int32, code int) error {
return c.SendAsync(ctx, proto.MessageFromServer, &mt.BadMsgNotification{
BadMsgID: badMsgID,
BadMsgSeqno: int(badSeqno),
ErrorCode: code,
})
}
// sendBadServerSalt 通知客户端修正 server_salterror_code 48
func (s *Server) sendBadServerSalt(ctx context.Context, c *Conn, badMsgID int64, badSeqno int32, newSalt int64) error {
return c.SendPriority(ctx, proto.MessageFromServer, &mt.BadServerSalt{
BadMsgID: badMsgID,
BadMsgSeqno: int(badSeqno),
ErrorCode: 48,
NewServerSalt: newSalt,
})
}
// typeName 返回 TL TypeID 的可读名称,未知时回退到 hex。
func (s *Server) typeName(id uint32) string {
if name := s.types.Get(id); name != "" {
return name
}
return fmt.Sprintf("%#x", id)
}
func validateClientEnvelope(now time.Time, msgID int64, seqNo int32, typeID uint32) int {
if msgID == 0 || proto.MessageID(msgID).Type() != proto.MessageFromClient {
return badMsgIDInvalidBits
}
msgTime := proto.MessageID(msgID).Time()
if msgTime.Before(now.Add(-300 * time.Second)) {
return badMsgIDTooLow
}
if msgTime.After(now.Add(30 * time.Second)) {
return badMsgIDTooHigh
}
if clientMessageAllowsEitherSeqParity(typeID) {
return 0
}
if clientMessageNeedsAck(typeID) {
if seqNo%2 == 0 {
return badMsgSeqNotOdd
}
} else if seqNo%2 != 0 {
return badMsgSeqNotEven
}
return 0
}
func validateClientContainer(containerMsgID int64, containerSeqNo int32, container proto.MessageContainer) int {
for _, m := range container.Messages {
if m.ID >= containerMsgID || int32(m.SeqNo) > containerSeqNo {
return badMsgContainer
}
typeID, err := (&bin.Buffer{Buf: m.Body}).PeekID()
if err != nil {
return badMsgContainer
}
if typeID == proto.MessageContainerTypeID {
return badMsgContainer
}
if code := validateClientContainerEnvelope(m.ID, int32(m.SeqNo), typeID); code != 0 {
return badMsgContainer
}
}
return 0
}
func validateClientContainerEnvelope(msgID int64, seqNo int32, typeID uint32) int {
if msgID == 0 || proto.MessageID(msgID).Type() != proto.MessageFromClient {
return badMsgIDInvalidBits
}
if clientMessageAllowsEitherSeqParity(typeID) {
return 0
}
if clientMessageNeedsAck(typeID) {
if seqNo%2 == 0 {
return badMsgSeqNotOdd
}
} else if seqNo%2 != 0 {
return badMsgSeqNotEven
}
return 0
}
func clientMessageAllowsEitherSeqParity(typeID uint32) bool {
switch typeID {
case mt.PingDelayDisconnectRequestTypeID,
// get_future_salts 的 seqno 奇偶在客户端间不一致:部分客户端按内容消息发奇数,
// gotd 按服务消息发偶数。两者都合法(官方服务器都接受),故不在此卡奇偶,避免
// 误判 bad_msg 触发客户端重连风暴。ack/content 行为仍由 clientMessageNeedsAck 决定。
mt.GetFutureSaltsRequestTypeID:
return true
default:
return false
}
}
func clientMessageNeedsAck(typeID uint32) bool {
switch typeID {
case proto.MessageContainerTypeID,
mt.MsgsAckTypeID,
mt.PingDelayDisconnectRequestTypeID,
mt.DestroySessionRequestTypeID,
mt.HTTPWaitRequestTypeID,
mt.BadMsgNotificationTypeID,
mt.BadServerSaltTypeID,
mt.MsgsAllInfoTypeID,
mt.MsgsStateInfoTypeID,
mt.MsgDetailedInfoTypeID,
mt.MsgNewDetailedInfoTypeID:
return false
default:
return true
}
}
func (cs *connState) seenRecord(msgID int64) (clientMsgRecord, bool) {
record, ok := cs.seen[msgID]
return record, ok
}
func (cs *connState) validateSeq(msgID int64, seqNo int32, content bool) int {
if !content {
return 0
}
// 快路径msg_id 与 seq_no 都严格高于已接受 content 高水位时,任何已见记录都不可能
// 与本条构成 too_low/too_high 反转,免去 O(len(seen)) 全扫描(正常客户端恒命中)。
if msgID > cs.maxContentMsgID && seqNo > cs.maxContentSeqNo {
return 0
}
for seenMsgID, record := range cs.seen {
if !record.content {
continue
}
if seenMsgID < msgID && record.seqNo >= seqNo {
return badMsgSeqTooLow
}
if seenMsgID > msgID && record.seqNo <= seqNo {
return badMsgSeqTooHigh
}
}
return 0
}
func (cs *connState) track(msgID int64, seqNo int32, content bool, state byte) {
cs.seen[msgID] = clientMsgRecord{
state: state,
seqNo: seqNo,
content: content,
}
if content {
if msgID > cs.maxContentMsgID {
cs.maxContentMsgID = msgID
}
if seqNo > cs.maxContentSeqNo {
cs.maxContentSeqNo = seqNo
}
}
cs.order = append(cs.order, msgID)
if msgID < cs.minSeen {
cs.minSeen = msgID
}
if msgID > cs.maxSeen {
cs.maxSeen = msgID
}
if len(cs.order) > maxTrackedClientMsgIDs {
oldest := cs.order[0]
cs.order = cs.order[1:]
delete(cs.seen, oldest)
if oldest == cs.minSeen || oldest == cs.maxSeen {
cs.recomputeRange()
}
}
}
func (cs *connState) stateInfo(msgIDs []int64) []byte {
info := make([]byte, len(msgIDs))
if len(cs.seen) == 0 {
for i := range info {
info[i] = msgStateUnknown
}
return info
}
for i, id := range msgIDs {
if id < cs.minSeen {
info[i] = msgStateUnknown
continue
}
if id > cs.maxSeen {
info[i] = msgStateNotReceivedHigh
continue
}
record, ok := cs.seen[id]
if !ok {
info[i] = msgStateNotReceived
continue
}
info[i] = record.state
}
return info
}
func (cs *connState) recomputeRange() {
cs.minSeen = math.MaxInt64
cs.maxSeen = 0
for id := range cs.seen {
if id < cs.minSeen {
cs.minSeen = id
}
if id > cs.maxSeen {
cs.maxSeen = id
}
}
if len(cs.seen) == 0 {
cs.minSeen = math.MaxInt64
}
}