fix: sync scoped connection and outbox exclusion updates

This commit is contained in:
A 2026-07-12 12:01:53 +08:00
parent aa21bd04e1
commit cbccd6a8d9
58 changed files with 919 additions and 1435 deletions

View file

@ -1,13 +1,13 @@
// Package store 定义存储接口与协议层 DTO不含具体实现。
//
// 布局主包只放接口AuthKeyStore / SessionStore / UserStore / AuthorizationStore /
// 布局主包只放接口AuthKeyStore / UserStore / AuthorizationStore /
// CodeStore / UpdateStateStore / UpdateEventStore 等)与协议 DTO三种后端实现各自独立成对称子包
// - store/memory —— 内存实现,测试替身与本地兜底
// - store/postgres —— PostgreSQLpgx + sqlc 生成查询 + golang-migrate 迁移)
// - store/redisstore —— Redisgo-redis
//
// 类型边界:接口签名分两类——
// - 协议产物用 store 自有 DTOAuthKeyData、SessionData、PhoneCode不依赖 tg.*,也非业务实体);
// - 协议产物用 store 自有 DTOAuthKeyData、PhoneCode不依赖 tg.*,也非业务实体);
// - 业务实体直接用 domainUserStore / AuthorizationStore / MessageStore / UpdateEventStore
// 收发 domain.User / domain.Authorization / domain.Message / domain.UpdateEvent。
package store

View file

@ -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

View file

@ -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, `

View file

@ -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

View file

@ -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)
}
}

View 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)
}
})
}
}

View file

@ -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),

View file

@ -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),

View file

@ -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),

View file

@ -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),

View file

@ -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),

View file

@ -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),

View file

@ -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),

View file

@ -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),

View file

@ -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),

View file

@ -1,4 +1,4 @@
// Package redisstore 用 Redis 实现高频易失态的存储接口第一阶段SessionStore
// Package redisstore 用 Redis 实现高频、易失且可重建的短状态、缓存、计数器与限流
//
// 职责边界见 docs/persistence-layer.md §1Redis 存「态与计数」,丢失可由 PG/协议恢复。
package redisstore

View file

@ -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 中的序列化形态(不含 IDID 即 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
}

View file

@ -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")
}
}

View file

@ -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
}