fix: sync scoped connection and outbox exclusion updates
This commit is contained in:
parent
aa21bd04e1
commit
cbccd6a8d9
58 changed files with 919 additions and 1435 deletions
|
|
@ -1,13 +1,13 @@
|
|||
// Package store 定义存储接口与协议层 DTO,不含具体实现。
|
||||
//
|
||||
// 布局:主包只放接口(AuthKeyStore / SessionStore / UserStore / AuthorizationStore /
|
||||
// 布局:主包只放接口(AuthKeyStore / UserStore / AuthorizationStore /
|
||||
// CodeStore / UpdateStateStore / UpdateEventStore 等)与协议 DTO;三种后端实现各自独立成对称子包:
|
||||
// - store/memory —— 内存实现,测试替身与本地兜底
|
||||
// - store/postgres —— PostgreSQL(pgx + sqlc 生成查询 + golang-migrate 迁移)
|
||||
// - store/redisstore —— Redis(go-redis)
|
||||
//
|
||||
// 类型边界:接口签名分两类——
|
||||
// - 协议产物用 store 自有 DTO:AuthKeyData、SessionData、PhoneCode(不依赖 tg.*,也非业务实体);
|
||||
// - 协议产物用 store 自有 DTO:AuthKeyData、PhoneCode(不依赖 tg.*,也非业务实体);
|
||||
// - 业务实体直接用 domain:UserStore / AuthorizationStore / MessageStore / UpdateEventStore
|
||||
// 收发 domain.User / domain.Authorization / domain.Message / domain.UpdateEvent。
|
||||
package store
|
||||
|
|
|
|||
|
|
@ -73,38 +73,6 @@ func (s *AuthKeyStore) Delete(_ context.Context, id [8]byte) error {
|
|||
return nil
|
||||
}
|
||||
|
||||
// SessionStore 是 store.SessionStore 的内存实现。
|
||||
type SessionStore struct {
|
||||
mu sync.RWMutex
|
||||
sessions map[int64]store.SessionData
|
||||
}
|
||||
|
||||
// NewSessionStore 创建内存 SessionStore。
|
||||
func NewSessionStore() *SessionStore {
|
||||
return &SessionStore{sessions: make(map[int64]store.SessionData)}
|
||||
}
|
||||
|
||||
func (s *SessionStore) Save(_ context.Context, d store.SessionData) error {
|
||||
s.mu.Lock()
|
||||
s.sessions[d.ID] = d
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SessionStore) Get(_ context.Context, id int64) (store.SessionData, bool, error) {
|
||||
s.mu.RLock()
|
||||
d, ok := s.sessions[id]
|
||||
s.mu.RUnlock()
|
||||
return d, ok, nil
|
||||
}
|
||||
|
||||
func (s *SessionStore) Delete(_ context.Context, id int64) error {
|
||||
s.mu.Lock()
|
||||
delete(s.sessions, id)
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
// TempAuthKeyBindingStore 是 store.TempAuthKeyBindingStore 的内存实现。
|
||||
type TempAuthKeyBindingStore struct {
|
||||
mu sync.RWMutex
|
||||
|
|
|
|||
|
|
@ -169,7 +169,7 @@ func TestDispatchOutboxLifecycleKeepsDurableEvents(t *testing.T) {
|
|||
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: owner.ID + int64(pts)},
|
||||
Bool: pts%2 == 0,
|
||||
}
|
||||
if _, err := events.AppendAllocatedWithDispatch(ctx, owner.ID, event, [8]byte{}, sessionID); err != nil {
|
||||
if _, err := events.AppendAllocatedWithDispatch(ctx, owner.ID, event, [8]byte{1}, sessionID); err != nil {
|
||||
t.Fatalf("AppendAllocatedWithDispatch pts=%d: %v", pts, err)
|
||||
}
|
||||
if _, err := tx.Exec(ctx, `
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ package postgres
|
|||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
|
|
@ -18,6 +19,21 @@ const (
|
|||
maxDispatchPoisonCleanupBatch = 1000
|
||||
)
|
||||
|
||||
var errInvalidDispatchOutboxExclusionPair = errors.New("dispatch outbox exclusion requires both raw auth key and session id")
|
||||
|
||||
// enqueueDispatch is the only production write boundary for dispatch_outbox.
|
||||
// A zero pair means no originating session is excluded; a non-zero pair identifies
|
||||
// one exact physical raw-auth/session tuple. A half pair is never meaningful because
|
||||
// session IDs are not globally unique and must fail the surrounding transaction.
|
||||
func enqueueDispatch(ctx context.Context, q *sqlcgen.Queries, arg sqlcgen.EnqueueDispatchParams) error {
|
||||
hasAuthKey := arg.ExcludeAuthKeyID != 0
|
||||
hasSession := arg.ExcludeSessionID != 0
|
||||
if hasAuthKey != hasSession {
|
||||
return errInvalidDispatchOutboxExclusionPair
|
||||
}
|
||||
return q.EnqueueDispatch(ctx, arg)
|
||||
}
|
||||
|
||||
// DispatchOutboxStore 用 PostgreSQL 实现 transactional outbox。
|
||||
type DispatchOutboxStore struct {
|
||||
q *sqlcgen.Queries
|
||||
|
|
|
|||
|
|
@ -0,0 +1,82 @@
|
|||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
|
||||
"telesrv/internal/domain"
|
||||
)
|
||||
|
||||
func TestDispatchOutboxExclusionPairInvariantPostgres(t *testing.T) {
|
||||
pool := testPool(t)
|
||||
ctx := context.Background()
|
||||
suffix := randomSuffix(t)
|
||||
owner := createTestUser(t, ctx, NewUserStore(pool), "+1887"+suffix+"01", "OutboxPair", "")
|
||||
t.Cleanup(func() { _, _ = pool.Exec(ctx, "DELETE FROM users WHERE id = $1", owner.ID) })
|
||||
|
||||
event := domain.UpdateEvent{
|
||||
Type: domain.UpdateEventDialogPinned,
|
||||
PtsCount: 1,
|
||||
Date: 1700002300,
|
||||
Peer: domain.Peer{Type: domain.PeerTypeUser, ID: owner.ID},
|
||||
Bool: true,
|
||||
}
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
authKeyID [8]byte
|
||||
sessionID int64
|
||||
}{
|
||||
{name: "auth key only", authKeyID: [8]byte{1}},
|
||||
{name: "session only", sessionID: 77},
|
||||
} {
|
||||
t.Run("write boundary "+test.name, func(t *testing.T) {
|
||||
_, err := NewUpdateEventStore(pool).AppendAllocatedWithDispatch(ctx, owner.ID, event, test.authKeyID, test.sessionID)
|
||||
if !errors.Is(err, errInvalidDispatchOutboxExclusionPair) {
|
||||
t.Fatalf("AppendAllocatedWithDispatch error = %v, want %v", err, errInvalidDispatchOutboxExclusionPair)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
var eventCount int
|
||||
if err := pool.QueryRow(ctx, "SELECT count(*)::int FROM user_update_events WHERE user_id = $1", owner.ID).Scan(&eventCount); err != nil {
|
||||
t.Fatalf("count events after rejected writes: %v", err)
|
||||
}
|
||||
if eventCount != 0 {
|
||||
t.Fatalf("events after rejected writes = %d, want 0 (transaction rollback)", eventCount)
|
||||
}
|
||||
|
||||
stored, err := NewUpdateEventStore(pool).AppendAllocated(ctx, owner.ID, event)
|
||||
if err != nil {
|
||||
t.Fatalf("append durable event for constraint test: %v", err)
|
||||
}
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
authKeyID int64
|
||||
sessionID int64
|
||||
}{
|
||||
{name: "auth key only", authKeyID: 1},
|
||||
{name: "session only", sessionID: 77},
|
||||
} {
|
||||
t.Run("database constraint "+test.name, func(t *testing.T) {
|
||||
_, err := pool.Exec(ctx, `
|
||||
INSERT INTO dispatch_outbox (
|
||||
target_user_id, pts, event_type, exclude_auth_key_id, exclude_session_id
|
||||
) VALUES ($1, $2, $3, $4, $5)`, owner.ID, stored.Pts, string(stored.Type), test.authKeyID, test.sessionID)
|
||||
var pgErr *pgconn.PgError
|
||||
if !errors.As(err, &pgErr) || pgErr.Code != "23514" || pgErr.ConstraintName != "dispatch_outbox_exclusion_pair_check" {
|
||||
t.Fatalf("direct insert error = %v, want check violation from dispatch_outbox_exclusion_pair_check", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
var outboxCount int
|
||||
if err := pool.QueryRow(ctx, "SELECT count(*)::int FROM dispatch_outbox WHERE target_user_id = $1", owner.ID).Scan(&outboxCount); err != nil {
|
||||
t.Fatalf("count outbox after rejected inserts: %v", err)
|
||||
}
|
||||
if outboxCount != 0 {
|
||||
t.Fatalf("outbox rows after rejected inserts = %d, want 0", outboxCount)
|
||||
}
|
||||
}
|
||||
31
internal/store/postgres/dispatch_outbox_exclusion_test.go
Normal file
31
internal/store/postgres/dispatch_outbox_exclusion_test.go
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
package postgres
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"telesrv/internal/store/postgres/sqlcgen"
|
||||
)
|
||||
|
||||
func TestEnqueueDispatchRejectsHalfExclusionPair(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
authKeyID int64
|
||||
sessionID int64
|
||||
}{
|
||||
{name: "auth key only", authKeyID: 1},
|
||||
{name: "session only", sessionID: 1},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
err := enqueueDispatch(context.Background(), nil, sqlcgen.EnqueueDispatchParams{
|
||||
ExcludeAuthKeyID: test.authKeyID,
|
||||
ExcludeSessionID: test.sessionID,
|
||||
})
|
||||
if !errors.Is(err, errInvalidDispatchOutboxExclusionPair) {
|
||||
t.Fatalf("enqueueDispatch error = %v, want %v", err, errInvalidDispatchOutboxExclusionPair)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -186,7 +186,7 @@ func (s *MessageStore) DeliverLoginCodeMessage(ctx context.Context, req domain.L
|
|||
if err := appendNewMessageEvent(ctx, qtx, msg); err != nil {
|
||||
return domain.LoginCodeDeliveryResult{}, err
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: req.UserID,
|
||||
Pts: int32(msg.Pts),
|
||||
EventType: string(domain.UpdateEventNewMessage),
|
||||
|
|
|
|||
|
|
@ -227,7 +227,7 @@ WHERE sender_user_id = $1
|
|||
dispatchAuthKeyID = excludeAuthKeyID
|
||||
dispatchSessionID = excludeSessionID
|
||||
}
|
||||
if err := q.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, q, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: userID,
|
||||
Pts: int32(deletePts),
|
||||
EventType: string(domain.UpdateEventDeleteMessages),
|
||||
|
|
@ -256,7 +256,7 @@ WHERE sender_user_id = $1
|
|||
}); err != nil {
|
||||
return res, fmt.Errorf("advance dialog read inbox after delete correction: %w", err)
|
||||
}
|
||||
if err := q.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, q, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: userID,
|
||||
Pts: int32(correction.Pts),
|
||||
EventType: string(domain.UpdateEventReadHistoryInbox),
|
||||
|
|
|
|||
|
|
@ -161,7 +161,7 @@ WHERE owner_user_id = $1 AND box_id = $2`, box.OwnerUserID, box.BoxID, int32(pts
|
|||
if err := appendUserUpdateEvent(ctx, tx, qtx, msg.OwnerUserID, event); err != nil {
|
||||
return res, fmt.Errorf("append web page event: %w", err)
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: msg.OwnerUserID,
|
||||
Pts: int32(pts),
|
||||
EventType: string(domain.UpdateEventWebPage),
|
||||
|
|
@ -257,7 +257,7 @@ WHERE message_sender_id = $1 AND private_message_id = $2`, messageSenderID, targ
|
|||
dispatchAuthKeyID = req.OriginAuthKeyID
|
||||
dispatchSessionID = req.OriginSessionID
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: msg.OwnerUserID,
|
||||
Pts: int32(pts),
|
||||
EventType: string(domain.UpdateEventEditMessage),
|
||||
|
|
|
|||
|
|
@ -379,7 +379,7 @@ func (s *MessageStore) ReadHistory(ctx context.Context, req domain.ReadHistoryRe
|
|||
if err := appendUserUpdateEvent(ctx, tx, qtx, req.OwnerUserID, res.InboxEvent); err != nil {
|
||||
return res, fmt.Errorf("append read inbox event: %w", err)
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: req.OwnerUserID,
|
||||
Pts: int32(readerPts),
|
||||
EventType: string(domain.UpdateEventReadHistoryInbox),
|
||||
|
|
@ -414,7 +414,7 @@ func (s *MessageStore) ReadHistory(ctx context.Context, req domain.ReadHistoryRe
|
|||
if err := appendUserUpdateEvent(ctx, tx, qtx, candidate.SenderOwnerUserID, res.OutboxEvent); err != nil {
|
||||
return res, fmt.Errorf("append read outbox event: %w", err)
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: candidate.SenderOwnerUserID,
|
||||
Pts: int32(senderPts),
|
||||
EventType: string(domain.UpdateEventReadHistoryOutbox),
|
||||
|
|
|
|||
|
|
@ -141,7 +141,7 @@ func (s *MessageStore) PinPrivateMessage(ctx context.Context, req domain.PinPriv
|
|||
dispatchAuthKeyID = req.OriginAuthKeyID
|
||||
dispatchSessionID = req.OriginSessionID
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: side.userID,
|
||||
Pts: int32(pts),
|
||||
EventType: string(domain.UpdateEventPinnedMessages),
|
||||
|
|
@ -286,7 +286,7 @@ func (s *MessageStore) UnpinAllPrivateMessages(ctx context.Context, req domain.U
|
|||
dispatchAuthKeyID = req.OriginAuthKeyID
|
||||
dispatchSessionID = req.OriginSessionID
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: side.userID,
|
||||
Pts: int32(pts),
|
||||
EventType: string(domain.UpdateEventPinnedMessages),
|
||||
|
|
|
|||
|
|
@ -182,7 +182,7 @@ WHERE d.user_id = $1
|
|||
if err := appendUserUpdateEvent(ctx, tx, qtx, req.OwnerUserID, res.Event); err != nil {
|
||||
return res, fmt.Errorf("append read message contents event: %w", err)
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: req.OwnerUserID,
|
||||
Pts: int32(pts),
|
||||
EventType: string(domain.UpdateEventReadMessageContents),
|
||||
|
|
@ -242,7 +242,7 @@ RETURNING box_id`, senderID, senderPrivateMessageIDs[senderID])
|
|||
if err := appendUserUpdateEvent(ctx, tx, qtx, senderID, event); err != nil {
|
||||
return res, fmt.Errorf("append sender content read event: %w", err)
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: senderID,
|
||||
Pts: int32(senderPts),
|
||||
EventType: string(domain.UpdateEventReadMessageContents),
|
||||
|
|
|
|||
|
|
@ -298,7 +298,7 @@ func (s *MessageStore) sendPrivateTextOnce(ctx context.Context, req domain.SendP
|
|||
if err := appendNewMessageEvent(ctx, qtx, sender); err != nil {
|
||||
return domain.SendPrivateTextResult{}, err
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: req.SenderUserID,
|
||||
Pts: int32(senderPts),
|
||||
EventType: string(domain.UpdateEventNewMessage),
|
||||
|
|
@ -360,7 +360,7 @@ func (s *MessageStore) sendPrivateTextOnce(ctx context.Context, req domain.SendP
|
|||
if err := appendNewMessageEvent(ctx, qtx, recipient); err != nil {
|
||||
return domain.SendPrivateTextResult{}, err
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: req.RecipientUserID,
|
||||
Pts: int32(recipientPts),
|
||||
EventType: string(domain.UpdateEventNewMessage),
|
||||
|
|
|
|||
|
|
@ -92,7 +92,7 @@ func (s *PhoneChangeStore) ChangePhone(ctx context.Context, req domain.PhoneChan
|
|||
if err := appendUserUpdateEvent(ctx, tx, qtx, req.UserID, event); err != nil {
|
||||
return domain.PhoneChangeResult{}, fmt.Errorf("append phone change event: %w", err)
|
||||
}
|
||||
if err := qtx.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, qtx, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: req.UserID,
|
||||
Pts: int32(event.Pts),
|
||||
EventType: string(event.Type),
|
||||
|
|
|
|||
|
|
@ -122,7 +122,7 @@ func (s *UpdateEventStore) appendInTx(ctx context.Context, db sqlcgen.DBTX, q *s
|
|||
return domain.UpdateEvent{}, fmt.Errorf("append update event: %w", err)
|
||||
}
|
||||
if dispatch {
|
||||
if err := q.EnqueueDispatch(ctx, sqlcgen.EnqueueDispatchParams{
|
||||
if err := enqueueDispatch(ctx, q, sqlcgen.EnqueueDispatchParams{
|
||||
TargetUserID: userID,
|
||||
Pts: int32(event.Pts),
|
||||
EventType: string(event.Type),
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
// Package redisstore 用 Redis 实现高频易失态的存储接口(第一阶段:SessionStore)。
|
||||
// Package redisstore 用 Redis 实现高频、易失且可重建的短状态、缓存、计数器与限流。
|
||||
//
|
||||
// 职责边界见 docs/persistence-layer.md §1:Redis 存「态与计数」,丢失可由 PG/协议恢复。
|
||||
package redisstore
|
||||
|
|
|
|||
|
|
@ -1,76 +0,0 @@
|
|||
package redisstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
|
||||
"telesrv/internal/store"
|
||||
)
|
||||
|
||||
// DefaultSessionTTL 是 session 记录的默认过期时间。
|
||||
// session 是连接态:过期或丢失后,客户端重连会触发 new_session_created / bad_server_salt 重建,
|
||||
// 因此 TTL 不必很长。每个随机 session_id 都落一条记录且断连不删,过长的 TTL
|
||||
// 只会堆积死 session(移动端每次重连一条)。7 天足够覆盖常规离线窗口。
|
||||
const DefaultSessionTTL = 7 * 24 * time.Hour
|
||||
|
||||
// SessionStore 用 Redis 实现 store.SessionStore。
|
||||
type SessionStore struct {
|
||||
c *redis.Client
|
||||
ttl time.Duration
|
||||
}
|
||||
|
||||
// NewSessionStore 创建 Redis SessionStore。ttl<=0 表示永不过期。
|
||||
func NewSessionStore(c *redis.Client, ttl time.Duration) *SessionStore {
|
||||
return &SessionStore{c: c, ttl: ttl}
|
||||
}
|
||||
|
||||
func sessionKey(id int64) string {
|
||||
return fmt.Sprintf("session:%d", id)
|
||||
}
|
||||
|
||||
// sessionValue 是 SessionData 在 Redis 中的序列化形态(不含 ID,ID 即 key)。
|
||||
type sessionValue struct {
|
||||
AuthKeyID [8]byte `json:"auth_key_id"`
|
||||
Salt int64 `json:"salt"`
|
||||
LastSeen int64 `json:"last_seen"`
|
||||
}
|
||||
|
||||
// Save 实现 store.SessionStore。
|
||||
func (s *SessionStore) Save(ctx context.Context, d store.SessionData) error {
|
||||
v, err := json.Marshal(sessionValue{AuthKeyID: d.AuthKeyID, Salt: d.Salt, LastSeen: d.LastSeen})
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal session: %w", err)
|
||||
}
|
||||
if err := s.c.Set(ctx, sessionKey(d.ID), v, s.ttl).Err(); err != nil {
|
||||
return fmt.Errorf("redis set session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get 实现 store.SessionStore。不存在时 found=false。
|
||||
func (s *SessionStore) Get(ctx context.Context, id int64) (store.SessionData, bool, error) {
|
||||
raw, err := s.c.Get(ctx, sessionKey(id)).Bytes()
|
||||
if err != nil {
|
||||
if errors.Is(err, redis.Nil) {
|
||||
return store.SessionData{}, false, nil
|
||||
}
|
||||
return store.SessionData{}, false, fmt.Errorf("redis get session: %w", err)
|
||||
}
|
||||
var v sessionValue
|
||||
if err := json.Unmarshal(raw, &v); err != nil {
|
||||
return store.SessionData{}, false, fmt.Errorf("unmarshal session: %w", err)
|
||||
}
|
||||
return store.SessionData{ID: id, AuthKeyID: v.AuthKeyID, Salt: v.Salt, LastSeen: v.LastSeen}, true, nil
|
||||
}
|
||||
|
||||
func (s *SessionStore) Delete(ctx context.Context, id int64) error {
|
||||
if err := s.c.Del(ctx, sessionKey(id)).Err(); err != nil {
|
||||
return fmt.Errorf("redis delete session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
|
@ -1,52 +0,0 @@
|
|||
package redisstore
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"telesrv/internal/store"
|
||||
)
|
||||
|
||||
// TestSessionStoreRoundTrip 验证 session 落 Redis 后能用全新 store 实例原样读回。
|
||||
// 未设 TELESRV_TEST_REDIS_ADDR 则跳过。
|
||||
func TestSessionStoreRoundTrip(t *testing.T) {
|
||||
addr := os.Getenv("TELESRV_TEST_REDIS_ADDR")
|
||||
if addr == "" {
|
||||
t.Skip("set TELESRV_TEST_REDIS_ADDR to run redis integration test")
|
||||
}
|
||||
ctx := context.Background()
|
||||
c, err := Open(ctx, addr, "", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("open: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = c.Close() })
|
||||
|
||||
want := store.SessionData{
|
||||
ID: 0x1234beef,
|
||||
AuthKeyID: [8]byte{1, 2, 3, 4, 5, 6, 7, 8},
|
||||
Salt: 42,
|
||||
LastSeen: 1000,
|
||||
}
|
||||
t.Cleanup(func() { _ = c.Del(ctx, sessionKey(want.ID)).Err() })
|
||||
|
||||
if err := NewSessionStore(c, time.Minute).Save(ctx, want); err != nil {
|
||||
t.Fatalf("save: %v", err)
|
||||
}
|
||||
|
||||
got, found, err := NewSessionStore(c, time.Minute).Get(ctx, want.ID)
|
||||
if err != nil {
|
||||
t.Fatalf("get: %v", err)
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("session not found after save")
|
||||
}
|
||||
if got != want {
|
||||
t.Fatalf("round trip mismatch: got %+v want %+v", got, want)
|
||||
}
|
||||
|
||||
if _, found, _ := NewSessionStore(c, time.Minute).Get(ctx, 999999); found {
|
||||
t.Fatal("unexpected found for missing session")
|
||||
}
|
||||
}
|
||||
|
|
@ -1,23 +0,0 @@
|
|||
package store
|
||||
|
||||
import "context"
|
||||
|
||||
// SessionData 是一条 MTProto session 记录(client 生成的 session_id)。
|
||||
//
|
||||
// 后续里程碑会扩展 device / layer 等字段。
|
||||
type SessionData struct {
|
||||
ID int64 // session_id(客户端生成)
|
||||
AuthKeyID [8]byte // 绑定的 auth key
|
||||
Salt int64 // 当前 server salt
|
||||
LastSeen int64 // unix 秒
|
||||
}
|
||||
|
||||
// SessionStore 记录在线 MTProto session。实现见 store/memory(测试替身)、store/redisstore。
|
||||
type SessionStore interface {
|
||||
// Save 保存或更新一条 session 记录。
|
||||
Save(ctx context.Context, s SessionData) error
|
||||
// Get 按 session_id 查询;不存在时 found=false。
|
||||
Get(ctx context.Context, id int64) (data SessionData, found bool, err error)
|
||||
// Delete 删除一条 session 记录;不存在时不报错。
|
||||
Delete(ctx context.Context, id int64) error
|
||||
}
|
||||
Loading…
Add table
Add a link
Reference in a new issue