diff --git a/internal/app/auth/service.go b/internal/app/auth/service.go index 643df5af..c9e52247 100644 --- a/internal/app/auth/service.go +++ b/internal/app/auth/service.go @@ -518,6 +518,13 @@ func (s *Service) Authorization(ctx context.Context, authKeyID [8]byte) (domain. return s.auths.ByAuthKey(ctx, authKeyID) } +func (s *Service) UpdateAuthorizationLayer(ctx context.Context, authKeyID [8]byte, layer int) error { + if s == nil || s.auths == nil || authKeyID == ([8]byte{}) || layer <= 0 { + return nil + } + return s.auths.UpdateLayer(ctx, authKeyID, layer) +} + func (s *Service) ListAuthorizations(ctx context.Context, userID int64) ([]domain.Authorization, error) { if s == nil || s.auths == nil || userID == 0 { return nil, nil diff --git a/internal/mtprotoedge/layer_downgrade_test.go b/internal/mtprotoedge/layer_downgrade_test.go index e258c3a6..730e22bb 100644 --- a/internal/mtprotoedge/layer_downgrade_test.go +++ b/internal/mtprotoedge/layer_downgrade_test.go @@ -2,10 +2,13 @@ package mtprotoedge import ( "bytes" + "encoding/binary" "testing" "github.com/gotd/td/bin" + "github.com/gotd/td/proto" "github.com/gotd/td/tg" + "go.uber.org/zap/zaptest" ) // TestConnDowngradedClone verifies the outbound seam downgrades a canonical @@ -63,3 +66,141 @@ func TestConnDowngradedClone(t *testing.T) { t.Errorf("227 passthrough altered bytes") } } + +func TestEncodeRPCResultDowngradesDifferenceMessagesForNegotiatedLayer225(t *testing.T) { + const ( + message227CRC = 0x7600b9d3 + message225CRC = 0x95ef6f2b + ) + c := &Conn{metrics: NopMetrics{}} + c.SetClientLayer(225) + diff := &tg.UpdatesDifference{ + NewMessages: []tg.MessageClass{ + &tg.Message{ + ID: 2, + FromID: &tg.PeerUser{UserID: 3}, + PeerID: &tg.PeerUser{UserID: 3}, + Date: 1, + Message: "hi", + }, + }, + NewEncryptedMessages: []tg.EncryptedMessageClass{}, + OtherUpdates: []tg.UpdateClass{}, + Chats: []tg.ChatClass{}, + Users: []tg.UserClass{}, + State: tg.UpdatesState{Pts: 2, Date: 1}, + } + + s := &Server{log: zaptest.NewLogger(t)} + encoded, err := s.encodeRPCResult(c, 12345, diff) + if err != nil { + t.Fatalf("encode rpc_result: %v", err) + } + var result proto.Result + if err := result.Decode(&bin.Buffer{Buf: encoded.body}); err != nil { + t.Fatalf("decode rpc_result: %v", err) + } + if result.RequestMessageID != 12345 { + t.Fatalf("req_msg_id = %d, want 12345", result.RequestMessageID) + } + if !bytes.Contains(result.Result, littleEndianID(message225CRC)) { + t.Fatalf("rpc_result inner object does not contain layer 225 message id %#08x", message225CRC) + } + if bytes.Contains(result.Result, littleEndianID(message227CRC)) { + t.Fatalf("rpc_result inner object still contains canonical message id %#08x", message227CRC) + } +} + +func TestEncodeRPCResultDowngradesDialogMessagesForNegotiatedLayer225(t *testing.T) { + const ( + message227CRC = 0x7600b9d3 + message225CRC = 0x95ef6f2b + ) + c := &Conn{metrics: NopMetrics{}} + c.SetClientLayer(225) + dialogs := &tg.MessagesDialogs{ + Dialogs: []tg.DialogClass{ + &tg.Dialog{ + Peer: &tg.PeerUser{UserID: 3}, + TopMessage: 2, + NotifySettings: tg.PeerNotifySettings{}, + }, + }, + Messages: []tg.MessageClass{ + &tg.Message{ + ID: 2, + FromID: &tg.PeerUser{UserID: 3}, + PeerID: &tg.PeerUser{UserID: 3}, + Date: 1, + Message: "hi", + }, + }, + Chats: []tg.ChatClass{}, + Users: []tg.UserClass{ + &tg.User{ID: 3, AccessHash: 5, FirstName: "A"}, + }, + } + + s := &Server{log: zaptest.NewLogger(t)} + encoded, err := s.encodeRPCResult(c, 12345, dialogs) + if err != nil { + t.Fatalf("encode rpc_result: %v", err) + } + var result proto.Result + if err := result.Decode(&bin.Buffer{Buf: encoded.body}); err != nil { + t.Fatalf("decode rpc_result: %v", err) + } + if !bytes.Contains(result.Result, littleEndianID(message225CRC)) { + t.Fatalf("rpc_result inner object does not contain layer 225 message id %#08x", message225CRC) + } + if bytes.Contains(result.Result, littleEndianID(message227CRC)) { + t.Fatalf("rpc_result inner object still contains canonical message id %#08x", message227CRC) + } +} + +func TestConnDowngradedCloneDowngradesUpdateNewMessageForLayer225(t *testing.T) { + const ( + message227CRC = 0x7600b9d3 + message225CRC = 0x95ef6f2b + ) + updates := &tg.Updates{ + Updates: []tg.UpdateClass{ + &tg.UpdateNewMessage{ + Message: &tg.Message{ + ID: 2, + FromID: &tg.PeerUser{UserID: 3}, + PeerID: &tg.PeerUser{UserID: 3}, + Date: 1, + Message: "hi", + }, + Pts: 2, + PtsCount: 1, + }, + }, + Users: []tg.UserClass{ + &tg.User{ID: 3, AccessHash: 5, FirstName: "A"}, + }, + Chats: []tg.ChatClass{}, + Date: 1, + Seq: 1, + } + enc, err := encodeOutboundMessage(updates) + if err != nil { + t.Fatalf("encode updates: %v", err) + } + c := &Conn{metrics: NopMetrics{}} + c.SetClientLayer(225) + out := c.downgradedClone(enc) + if !bytes.Contains(out.body, littleEndianID(message225CRC)) { + t.Fatalf("push update does not contain layer 225 message id %#08x", message225CRC) + } + if bytes.Contains(out.body, littleEndianID(message227CRC)) { + t.Fatalf("push update still contains canonical message id %#08x", message227CRC) + } +} + +func littleEndianID(id uint32) []byte { + buf := make([]byte, 4) + binary.LittleEndian.PutUint32(buf, id) + return buf +} diff --git a/internal/rpc/compat.go b/internal/rpc/compat.go index f870e228..59b6a8fe 100644 --- a/internal/rpc/compat.go +++ b/internal/rpc/compat.go @@ -2,18 +2,10 @@ package rpc import "context" -// withAndroidCompatMetadata 为「客户端构造器漂移」请求**仅在当前 ctx 内**兜底 client 层/类型。 -// 这类请求多来自未完整走 initConnection 的 DrKLO Android;当 layer/ClientType 仍未知时 -// 按 android 处理,使下游(createChat 的 legacy 响应、langpack 的 lang_pack 派生)行为正确。 -// -// 关键:**绝不把这个兜底值写回持久缓存**(不调 rememberClientLayer/rememberClientInfo)。 -// 缓存是 invokeWithLayer/initConnection 的权威产物;条目被驱逐时该兜底会拿不到真实值而误判 -// 227/android,若写回就把长连接老客户端的真实 layer/类型永久覆盖(与 NegotiatedLayer 的 -// 「驱逐时不覆盖」契约矛盾)。出站 layer 由 Conn.clientLayer 承载(非覆盖),与本兜底无关。 +// withAndroidCompatMetadata 为「客户端构造器漂移」请求仅兜底 client 类型。 +// DrKLO/OwpenGram Android 可能在不同版本使用不同 TL layer,client-private 构造器 +// 只能证明这是 Android 兼容路径,不能替代 invokeWithLayer 里的真实 layer。 func (r *Router) withAndroidCompatMetadata(ctx context.Context) context.Context { - if LayerFrom(ctx) == 0 { - ctx = WithLayer(ctx, currentClientLayer) - } if ClientTypeFrom(ctx) == ClientTypeUnknown { ctx = WithClientInfo(ctx, ClientInfo{LangPack: string(ClientTypeAndroid), Type: ClientTypeAndroid}) } diff --git a/internal/rpc/deps.go b/internal/rpc/deps.go index f4ec4148..0ab9334e 100644 --- a/internal/rpc/deps.go +++ b/internal/rpc/deps.go @@ -39,6 +39,7 @@ type AuthService interface { SignInBot(ctx context.Context, a domain.Authorization, token string) (domain.User, error) LogOut(ctx context.Context, authKeyID [8]byte) error Authorization(ctx context.Context, authKeyID [8]byte) (domain.Authorization, bool, error) + UpdateAuthorizationLayer(ctx context.Context, authKeyID [8]byte, layer int) error ListAuthorizations(ctx context.Context, userID int64) ([]domain.Authorization, error) ResetAuthorization(ctx context.Context, userID, hash int64) (domain.Authorization, bool, error) ResetAuthorizations(ctx context.Context, userID int64, keepAuthKeyID [8]byte) ([]domain.Authorization, error) diff --git a/internal/rpc/router.go b/internal/rpc/router.go index b1f8a74f..c3277c11 100644 --- a/internal/rpc/router.go +++ b/internal/rpc/router.go @@ -212,12 +212,13 @@ func (r *Router) Dispatch(ctx context.Context, authKeyID [8]byte, sessionID int6 ctx = WithUserID(ctx, userID) } tUser := r.clock.Now() - info, hasClientMetadata := r.clientSessionInfo(ctx) + info, hasClientMetadata, clientMetadataStored := r.clientSessionInfo(ctx) if hasUserID { if authInfo, ok := r.clientSessionInfoFromAuthorization(ctx, userID, effectiveAuthKeyID, info); ok { info = mergeClientSessionInfo(info, authInfo) hasClientMetadata = true r.rememberClientSessionInfo(ctx, info) + clientMetadataStored = true } } // 前置鉴权阶段(auth key 解析 / user 重校验 / client info)慢路径告警:超阈值才记,避免刷屏。 @@ -237,6 +238,9 @@ func (r *Router) Dispatch(ctx context.Context, authKeyID [8]byte, sessionID int6 } } if hasClientMetadata { + if !clientMetadataStored { + r.rememberClientSessionInfoIfMissing(ctx, info) + } if info.layer != 0 { ctx = WithLayer(ctx, info.layer) } @@ -590,8 +594,8 @@ func (r *Router) dispatch(ctx context.Context, b *bin.Buffer, depth int) (bin.En } id = newID if clientDrift { - // 客户端漂移多来自未完整 initConnection 的 DrKLO;按既有行为在 - // 类型/层未知时兜底为 android(withAndroidCompatMetadata 自带 unknown 守卫)。 + // 客户端漂移只能证明这是 Android 兼容路径;layer 仍以 + // invokeWithLayer 或授权记录里的真实观测值为准。 ctx = r.withAndroidCompatMetadata(ctx) } } @@ -675,9 +679,48 @@ func (r *Router) rememberClientInfo(ctx context.Context, info ClientInfo) { } func (r *Router) rememberClientLayer(ctx context.Context, layer int) { - r.mutateClientSessionInfo(ctx, func(sessionInfo *clientSessionInfo) { - sessionInfo.layer = layer - }) + if layer <= 0 { + return + } + rawAuthKeyID, ok := RawAuthKeyIDFrom(ctx) + if !ok { + return + } + sessionID, ok := SessionIDFrom(ctx) + if !ok { + return + } + authKeyID, hasAuthKeyID := AuthKeyIDFrom(ctx) + persistAuthLayer := false + r.clientInfoMu.Lock() + if r.clientInfo == nil { + r.clientInfo = make(map[clientInfoSessionKey]clientSessionInfo) + } + if hasAuthKeyID { + if info, ok := r.authInfo[authKeyID]; !ok || info.layer != layer { + persistAuthLayer = true + } + } + sessionKey := clientInfoSessionKey{rawAuthKeyID: rawAuthKeyID, sessionID: sessionID} + sessionInfo, exists := r.clientInfo[sessionKey] + sessionInfo.layer = layer + if !exists { + evictMapEntryIfFullLocked(r.clientInfo, maxClientInfoEntries) + } + r.clientInfo[sessionKey] = sessionInfo + r.rememberAuthClientLayerLocked(rawAuthKeyID, layer) + if hasAuthKeyID { + r.rememberAuthClientLayerLocked(authKeyID, layer) + } + r.clientInfoMu.Unlock() + if persistAuthLayer && r.deps.Auth != nil { + if err := r.deps.Auth.UpdateAuthorizationLayer(ctx, authKeyID, layer); err != nil { + r.log.Warn("update authorization layer failed", + zap.Int("layer", layer), + zap.String("auth_key_id", fmt.Sprintf("%x", authKeyID[:])), + zap.Error(err)) + } + } } // NegotiatedLayer returns the TL layer the given session negotiated via @@ -753,6 +796,71 @@ func (r *Router) rememberClientSessionInfo(ctx context.Context, sessionInfo clie } } +func (r *Router) rememberClientSessionInfoIfMissing(ctx context.Context, sessionInfo clientSessionInfo) bool { + rawAuthKeyID, ok := RawAuthKeyIDFrom(ctx) + if !ok { + return false + } + sessionID, ok := SessionIDFrom(ctx) + if !ok { + return false + } + authKeyID, hasAuthKeyID := AuthKeyIDFrom(ctx) + if r.clientSessionInfoStored(rawAuthKeyID, sessionID, authKeyID, hasAuthKeyID, sessionInfo) { + return false + } + r.clientInfoMu.Lock() + defer r.clientInfoMu.Unlock() + if r.clientSessionInfoStoredLocked(rawAuthKeyID, sessionID, authKeyID, hasAuthKeyID, sessionInfo) { + return false + } + if r.clientInfo == nil { + r.clientInfo = make(map[clientInfoSessionKey]clientSessionInfo) + } + sessionKey := clientInfoSessionKey{rawAuthKeyID: rawAuthKeyID, sessionID: sessionID} + if _, exists := r.clientInfo[sessionKey]; !exists { + evictMapEntryIfFullLocked(r.clientInfo, maxClientInfoEntries) + } + r.clientInfo[sessionKey] = mergeClientSessionInfo(r.clientInfo[sessionKey], sessionInfo) + r.rememberAuthClientInfoLocked(rawAuthKeyID, sessionInfo) + if hasAuthKeyID { + r.rememberAuthClientInfoLocked(authKeyID, sessionInfo) + } + return true +} + +func (r *Router) clientSessionInfoStored(rawAuthKeyID [8]byte, sessionID int64, authKeyID [8]byte, hasAuthKeyID bool, required clientSessionInfo) bool { + r.clientInfoMu.RLock() + defer r.clientInfoMu.RUnlock() + return r.clientSessionInfoStoredLocked(rawAuthKeyID, sessionID, authKeyID, hasAuthKeyID, required) +} + +func (r *Router) clientSessionInfoStoredLocked(rawAuthKeyID [8]byte, sessionID int64, authKeyID [8]byte, hasAuthKeyID bool, required clientSessionInfo) bool { + if !clientSessionInfoContains(r.clientInfo[clientInfoSessionKey{rawAuthKeyID: rawAuthKeyID, sessionID: sessionID}], required) { + return false + } + if !clientSessionInfoContains(r.authInfo[rawAuthKeyID], required) { + return false + } + if hasAuthKeyID && !clientSessionInfoContains(r.authInfo[authKeyID], required) { + return false + } + return true +} + +func clientSessionInfoContains(current, required clientSessionInfo) bool { + if required.layer != 0 && current.layer != required.layer { + return false + } + if required.hasClientInfo && (!current.hasClientInfo || current.clientInfo != required.clientInfo) { + return false + } + if required.authorizationChecked && !current.authorizationChecked { + return false + } + return true +} + func (r *Router) rememberAuthClientInfoLocked(authKeyID [8]byte, info clientSessionInfo) { if r.authInfo == nil { r.authInfo = make(map[[8]byte]clientSessionInfo) @@ -764,6 +872,21 @@ func (r *Router) rememberAuthClientInfoLocked(authKeyID [8]byte, info clientSess r.authInfo[authKeyID] = mergeClientSessionInfo(current, info) } +func (r *Router) rememberAuthClientLayerLocked(authKeyID [8]byte, layer int) { + if layer <= 0 { + return + } + if r.authInfo == nil { + r.authInfo = make(map[[8]byte]clientSessionInfo) + } + if _, exists := r.authInfo[authKeyID]; !exists { + evictMapEntryIfFullLocked(r.authInfo, maxAuthInfoEntries) + } + info := r.authInfo[authKeyID] + info.layer = layer + r.authInfo[authKeyID] = info +} + // forgetClientSessionInfo 随连接下线移除该 session 的元数据缓存条目,并清掉以该 raw // auth_key 为键的 authInfo 兜底条目,使 authInfo 收敛到活跃 raw auth key(主导的单 // session/key 场景下严格回收)。共享同一 raw auth_key 的其它 session 若仍在线,会在下一次 @@ -786,29 +909,33 @@ func evictMapEntryIfFullLocked[K comparable, V any](m map[K]V, limit int) { } } -func (r *Router) clientSessionInfo(ctx context.Context) (clientSessionInfo, bool) { +func (r *Router) clientSessionInfo(ctx context.Context) (clientSessionInfo, bool, bool) { rawAuthKeyID, ok := RawAuthKeyIDFrom(ctx) if !ok { - return clientSessionInfo{}, false + return clientSessionInfo{}, false, false } sessionID, ok := SessionIDFrom(ctx) if !ok { - return clientSessionInfo{}, false + return clientSessionInfo{}, false, false } r.clientInfoMu.RLock() defer r.clientInfoMu.RUnlock() info, ok := r.clientInfo[clientInfoSessionKey{rawAuthKeyID: rawAuthKeyID, sessionID: sessionID}] + authKeyID, hasAuthKeyID := AuthKeyIDFrom(ctx) if authInfo, authOK := r.authInfo[rawAuthKeyID]; authOK { info = mergeClientSessionInfo(info, authInfo) ok = true } - if authKeyID, hasAuthKeyID := AuthKeyIDFrom(ctx); hasAuthKeyID { + if hasAuthKeyID { if authInfo, authOK := r.authInfo[authKeyID]; authOK { info = mergeClientSessionInfo(info, authInfo) ok = true } } - return info, ok + if !ok { + return info, false, false + } + return info, true, r.clientSessionInfoStoredLocked(rawAuthKeyID, sessionID, authKeyID, hasAuthKeyID, info) } func (r *Router) cachedResolvedAuthClientInfo(authKeyID [8]byte) (clientSessionInfo, bool) { @@ -879,8 +1006,6 @@ func clientSessionInfoFromAuthorizationRecord(item domain.Authorization, current if info.layer == 0 { if current.layer != 0 { info.layer = current.layer - } else if info.clientInfo.ClientType() != ClientTypeUnknown { - info.layer = currentClientLayer } } return info diff --git a/internal/rpc/router_dispatch_test.go b/internal/rpc/router_dispatch_test.go index 163ee75e..5c2b84cf 100644 --- a/internal/rpc/router_dispatch_test.go +++ b/internal/rpc/router_dispatch_test.go @@ -133,7 +133,7 @@ func TestDispatchRemembersLayerAndClientTypeForSession(t *testing.T) { } sessionCtx := WithSessionID(WithRawAuthKeyID(WithAuthKeyID(context.Background(), rawAuthKeyID), rawAuthKeyID), sessionID) - info, ok := r.clientSessionInfo(sessionCtx) + info, ok, _ := r.clientSessionInfo(sessionCtx) if !ok { t.Fatalf("session metadata missing") } @@ -218,12 +218,15 @@ func TestAndroidLegacyCompatLogsClientMetadataWithoutInit(t *testing.T) { t.Fatalf("RPC inner handled log missing") } fields := entries[len(entries)-1].ContextMap() - if got := intLogField(fields["layer"]); got != currentClientLayer { - t.Fatalf("logged layer = %d fields=%v, want %d", got, fields, currentClientLayer) + if got := intLogField(fields["layer"]); got != 0 { + t.Fatalf("logged layer = %d fields=%v, want 0", got, fields) } if got := fields["client_type"]; got != string(ClientTypeAndroid) { t.Fatalf("logged client_type = %v, want %s", got, ClientTypeAndroid) } + if got, ok := r.NegotiatedLayer(rawAuthKeyID, sessionID); ok || got != currentClientLayer { + t.Fatalf("negotiated layer = (%d,%v), want (%d,false)", got, ok, currentClientLayer) + } } // TestNegotiatedLayerStickyContract pins the (layer, ok) contract that keeps the @@ -255,6 +258,67 @@ func TestNegotiatedLayerStickyContract(t *testing.T) { } } +func TestObservedClientLayerOverridesStaleAuthFallback(t *testing.T) { + r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{}, zaptest.NewLogger(t), clock.System) + rawAuthKey := [8]byte{0x68, 0x25, 0x7a, 0x01} + effectiveAuthKey := [8]byte{0x8b, 0x6f, 0x26, 0x17} + ctx := WithAuthKeyID( + WithSessionID(WithRawAuthKeyID(context.Background(), rawAuthKey), 100), + effectiveAuthKey, + ) + + r.rememberClientSessionInfo(ctx, clientSessionInfo{layer: currentClientLayer}) + r.rememberClientLayer(ctx, 225) + + if got, ok := r.NegotiatedLayer(rawAuthKey, 100); !ok || got != 225 { + t.Fatalf("exact session layer = (%d,%v), want (225,true)", got, ok) + } + if got, ok := r.NegotiatedLayer(rawAuthKey, 101); !ok || got != 225 { + t.Fatalf("raw auth fallback layer = (%d,%v), want (225,true)", got, ok) + } + if got, ok := r.NegotiatedLayer(effectiveAuthKey, 101); !ok || got != 225 { + t.Fatalf("effective auth fallback layer = (%d,%v), want (225,true)", got, ok) + } +} + +func TestInvokeWithLayerPersistsClientLayerUpgrade(t *testing.T) { + authKeyID := [8]byte{0x68, 0x25, 0x7a, 0x02} + userID := int64(1780269504) + auth := &captureAuthService{ + userID: userID, + authorizations: []domain.Authorization{{ + AuthKeyID: authKeyID, + UserID: userID, + Layer: 225, + }}, + } + r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{ + Auth: auth, + }, zaptest.NewLogger(t), clock.System) + req := &tg.InvokeWithLayerRequest{ + Layer: currentClientLayer, + Query: &tg.HelpGetConfigRequest{}, + } + var in bin.Buffer + if err := req.Encode(&in); err != nil { + t.Fatalf("encode invokeWithLayer: %v", err) + } + + if _, err := r.Dispatch(context.Background(), authKeyID, 100, &in); err != nil { + t.Fatalf("dispatch invokeWithLayer: %v", err) + } + + if auth.layerUpdates != 1 { + t.Fatalf("layer update calls = %d, want 1", auth.layerUpdates) + } + if got := auth.authorizations[0].Layer; got != currentClientLayer { + t.Fatalf("persisted layer = %d, want %d", got, currentClientLayer) + } + if got, ok := r.NegotiatedLayer(authKeyID, 101); !ok || got != currentClientLayer { + t.Fatalf("new session negotiated layer = (%d,%v), want (%d,true)", got, ok, currentClientLayer) + } +} + func TestClientTypeDetectsAndroidSDKVersion(t *testing.T) { info := normalizeClientInfo(ClientInfo{ DeviceModel: "GooglePixel 9a", @@ -331,6 +395,69 @@ func TestDispatchRestoresClientMetadataFromAuthorization(t *testing.T) { } } +func TestDispatchCopiesEffectiveAuthLayerToRawTempSession(t *testing.T) { + rawAuthKeyID := [8]byte{0x1a, 0x2d, 0x2d, 0x3d, 0x4b, 0x38, 0x62, 0xc0} + permAuthKeyID := [8]byte{0x5b, 0x1c, 0x12, 0x24, 0x98, 0x85, 0x60, 0xc1} + userID := int64(1780243218) + auth := &captureAuthService{ + resolvedAuthKeyID: permAuthKeyID, + hasResolved: true, + userID: userID, + authorizations: []domain.Authorization{{ + AuthKeyID: permAuthKeyID, + UserID: userID, + Layer: 225, + DeviceModel: "nubiaNX629J", + Platform: string(ClientTypeAndroid), + SystemVersion: "SDK 30", + AppVersion: "12.7.3 (67509) pbeta", + }}, + } + r := New(Config{DC: 2, IP: "127.0.0.1", Port: 2398}, Deps{ + Auth: auth, + }, zaptest.NewLogger(t), clock.System) + + var warm bin.Buffer + if err := (&tg.HelpGetConfigRequest{}).Encode(&warm); err != nil { + t.Fatalf("encode warm request: %v", err) + } + if _, err := r.Dispatch(context.Background(), permAuthKeyID, 11, &warm); err != nil { + t.Fatalf("dispatch warm perm request: %v", err) + } + if got, ok := r.NegotiatedLayer(permAuthKeyID, 12); !ok || got != 225 { + t.Fatalf("perm auth fallback layer = (%d,%v), want (225,true)", got, ok) + } + + var firstTempRequest bin.Buffer + if err := (&tg.HelpGetConfigRequest{}).Encode(&firstTempRequest); err != nil { + t.Fatalf("encode temp request: %v", err) + } + const tempSessionID = int64(5579025282411299519) + if _, err := r.Dispatch(context.Background(), rawAuthKeyID, tempSessionID, &firstTempRequest); err != nil { + t.Fatalf("dispatch first temp request: %v", err) + } + if got, ok := r.NegotiatedLayer(rawAuthKeyID, tempSessionID); !ok || got != 225 { + t.Fatalf("raw temp exact session layer = (%d,%v), want (225,true)", got, ok) + } + if got, ok := r.NegotiatedLayer(rawAuthKeyID, tempSessionID+1); !ok || got != 225 { + t.Fatalf("raw temp auth fallback layer = (%d,%v), want (225,true)", got, ok) + } + hotCtx := WithAuthKeyID( + WithSessionID(WithRawAuthKeyID(context.Background(), rawAuthKeyID), tempSessionID), + permAuthKeyID, + ) + info, ok, stored := r.clientSessionInfo(hotCtx) + if !ok || info.layer != 225 { + t.Fatalf("cached temp session info = (%+v,%v), want layer 225", info, ok) + } + if !stored { + t.Fatalf("cached temp session metadata is not fully materialized for hot path") + } + if wrote := r.rememberClientSessionInfoIfMissing(hotCtx, info); wrote { + t.Fatalf("hot path rewrote already materialized client session metadata") + } +} + func TestDispatchRestoresAndroidMetadataFromAuthorizationSDKVersion(t *testing.T) { core, logs := observer.New(zap.DebugLevel) authKeyID := [8]byte{0x16, 0x65, 0x54, 0x12, 0xaa, 0xbb, 0xcc, 0xdd} @@ -363,8 +490,8 @@ func TestDispatchRestoresAndroidMetadataFromAuthorizationSDKVersion(t *testing.T t.Fatalf("RPC inner handled log missing") } fields := entries[len(entries)-1].ContextMap() - if got := intLogField(fields["layer"]); got != currentClientLayer { - t.Fatalf("logged layer = %d fields=%v, want %d", got, fields, currentClientLayer) + if got := intLogField(fields["layer"]); got != 0 { + t.Fatalf("logged layer = %d fields=%v, want 0", got, fields) } if got := fields["client_type"]; got != string(ClientTypeAndroid) { t.Fatalf("logged client_type = %v, want %s", got, ClientTypeAndroid) diff --git a/internal/rpc/rpc_testkit_auth_test.go b/internal/rpc/rpc_testkit_auth_test.go index 5f662ec0..9735e59e 100644 --- a/internal/rpc/rpc_testkit_auth_test.go +++ b/internal/rpc/rpc_testkit_auth_test.go @@ -24,6 +24,7 @@ type captureAuthService struct { authorizations []domain.Authorization authorizationLookups int authorizationLists int + layerUpdates int loggedOutAuthKeyID [8]byte pendingPasswordUserID int64 pendingPassword bool @@ -111,6 +112,10 @@ func (s *blockingUserAuthService) Authorization(context.Context, [8]byte) (domai return domain.Authorization{}, false, nil } +func (s *blockingUserAuthService) UpdateAuthorizationLayer(context.Context, [8]byte, int) error { + return nil +} + func (s *blockingUserAuthService) ListAuthorizations(context.Context, int64) ([]domain.Authorization, error) { return nil, nil } @@ -222,6 +227,17 @@ func (s *captureAuthService) Authorization(_ context.Context, authKeyID [8]byte) return domain.Authorization{}, false, nil } +func (s *captureAuthService) UpdateAuthorizationLayer(_ context.Context, authKeyID [8]byte, layer int) error { + s.layerUpdates++ + for i := range s.authorizations { + if s.authorizations[i].AuthKeyID == authKeyID { + s.authorizations[i].Layer = layer + return nil + } + } + return nil +} + func (s *captureAuthService) ListAuthorizations(context.Context, int64) ([]domain.Authorization, error) { s.authorizationLists++ return append([]domain.Authorization(nil), s.authorizations...), nil diff --git a/internal/store/authorization.go b/internal/store/authorization.go index 548d0a02..08e35552 100644 --- a/internal/store/authorization.go +++ b/internal/store/authorization.go @@ -10,6 +10,7 @@ import ( type AuthorizationStore interface { Bind(ctx context.Context, a domain.Authorization) error ByAuthKey(ctx context.Context, authKeyID [8]byte) (domain.Authorization, bool, error) + UpdateLayer(ctx context.Context, authKeyID [8]byte, layer int) error ListByUser(ctx context.Context, userID int64) ([]domain.Authorization, error) Delete(ctx context.Context, authKeyID [8]byte) error DeleteByHash(ctx context.Context, userID, hash int64) (domain.Authorization, bool, error) diff --git a/internal/store/memory/auth.go b/internal/store/memory/auth.go index 01fe3f38..87514b31 100644 --- a/internal/store/memory/auth.go +++ b/internal/store/memory/auth.go @@ -158,6 +158,20 @@ func (s *AuthorizationStore) ByAuthKey(_ context.Context, id [8]byte) (domain.Au return a, ok, nil } +func (s *AuthorizationStore) UpdateLayer(_ context.Context, id [8]byte, layer int) error { + if layer <= 0 { + return nil + } + s.mu.Lock() + if a, ok := s.m[id]; ok { + a.Layer = layer + a.ActiveAt = time.Now() + s.m[id] = a + } + s.mu.Unlock() + return nil +} + func (s *AuthorizationStore) MarkPasswordPassed(_ context.Context, id [8]byte) error { s.mu.Lock() if a, ok := s.m[id]; ok { diff --git a/internal/store/postgres/authorization.go b/internal/store/postgres/authorization.go index 1cc310dd..819cd2cd 100644 --- a/internal/store/postgres/authorization.go +++ b/internal/store/postgres/authorization.go @@ -68,6 +68,18 @@ FROM authorizations WHERE auth_key_id = $1`, authKeyIDToInt64(id)) return a, true, nil } +func (s *AuthorizationStore) UpdateLayer(ctx context.Context, id [8]byte, layer int) error { + if layer <= 0 { + return nil + } + if _, err := s.db.Exec(ctx, ` +UPDATE authorizations SET layer = $2, active_at = now() WHERE auth_key_id = $1`, + authKeyIDToInt64(id), int32(layer)); err != nil { + return fmt.Errorf("update authorization layer: %w", err) + } + return nil +} + // MarkPasswordPassed 在两步验证通过后清除 password_pending,使 auth_key 转为完全授权。 func (s *AuthorizationStore) MarkPasswordPassed(ctx context.Context, id [8]byte) error { if _, err := s.db.Exec(ctx, `