feat: 实现消息提交防重配额授权与最小分发

This commit is contained in:
Nixevol
2026-09-30 06:57:57 +08:00
parent 532ee44da3
commit 8bcb5c63be
8 changed files with 1242 additions and 26 deletions
+29 -1
View File
@@ -278,7 +278,35 @@
## 消息 M
暂无。
### M1 2026-09-30
1. **提交时分发做成最小正确版**
- 原条款:DEVELOPMENT 7.3 步骤 8 / 7.4:`send_at` 已到则同一写操作内完整分发(停用拒绝、`queue_full`、`expire_at`/宽限、无接收者 `completed`、回执等)。
- 实际做法:单聊只插一条 `pending`;群按当时 `group_members` 去掉发送者各插 `pending`;消息改为 `dispatched`。不设 `expire_at`,不检查接收端配额/在线/停用,不因无接收者改为 `completed`,不写回执,不唤醒推送循环。
- 原因:M1 范围是提交;完整分发与推送属 M2。
- 备选方案:M1 直接实现完整 7.4(抢 M2)。
- 影响:到点消息已有投递行,但停用成员仍会有 `pending`;无成员群仍为 `dispatched` 且无投递;推送需等 M2。
2. **请求频率突发容量写死为 100**
- 原条款:DEVELOPMENT 6.10 每端每秒 50、突发 100;配置示例仅有 `requests_per_second`。
- 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。速率桶挂在 `message.App` 的 `Submit` 入口;`ack`/`receipt_ack` 尚未实现故未接桶。
- 原因:配置无独立 burst 字段。
- 备选方案:配置增加 `request_burst`;由连接线在上行统一限流。
- 影响:改 `requests_per_second` 不改突发;正式接线后若 N 线也限流可能双重计数。
3. **未接线 `cmd/nixmsg`**
- 原条款:可替换 T0.4 假实现。
- 实际做法:新增 `message.App` 实现 `Submit`;保留 `Stub`;按任务隔离要求未改 `cmd/nixmsg`/`wire.go`。
- 原因:本任务禁止改 `cmd/nixmsg`;总控接线或后续任务再换。
- 备选方案:本任务直接改 `wire.go`。
- 影响:进程内仍用 Stub,需显式构造 `message.New` 才能用真实提交。
4. **防重键在、消息行已删时返回 `not_found`**
- 原条款:防重命中返回原消息当前状态;未写明消息行已被清理时的提交重试行为(状态查询为 `not_found`)。
- 实际做法:`send_keys` 指纹相同但 `messages` 无行时返回 `not_found`。
- 原因:无法构造 `send_at`/`state`。
- 备选方案:在 `send_keys` 冗余存结果快照。
- 影响:保留期过后的重试不再幂等成功。
## 身份 I
+153
View File
@@ -0,0 +1,153 @@
package message
import (
"context"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/config"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
// 消息状态(DEVELOPMENT 7.1)。
const (
StateScheduled = "scheduled"
StateDispatched = "dispatched"
StateCompleted = "completed"
)
// 投递状态。
const (
DeliveryPending = "pending"
)
// talk_grants.kind。
const (
GrantKindPassword = "password"
GrantKindReply = "reply"
)
// 请求频率桶默认突发容量(DEVELOPMENT 6.10;配置无单独字段)。
const defaultRequestBurst = 100
// Limits 是提交所需的配置上限(来自 config.LimitsConfig)。
type Limits struct {
MaxBodyBytes int
MaxMetaBytes int
MaxFrameBytes int
MaxTTLSeconds int64
MaxScheduleSeconds int64
RequestsPerSecond float64
RequestBurst int
MaxPendingPerSender int
MaxPendingPerReceiver int
GraceSeconds int64
}
// LimitsFromConfig 从平台配置构造 Limits。
func LimitsFromConfig(c config.LimitsConfig) Limits {
burst := defaultRequestBurst
return Limits{
MaxBodyBytes: c.MaxBodyBytes,
MaxMetaBytes: c.MaxMetaBytes,
MaxFrameBytes: c.MaxFrameBytes,
MaxTTLSeconds: int64(c.MaxTTLSeconds),
MaxScheduleSeconds: int64(c.MaxScheduleSeconds),
RequestsPerSecond: float64(c.RequestsPerSecond),
RequestBurst: burst,
MaxPendingPerSender: c.MaxPendingPerSender,
MaxPendingPerReceiver: c.MaxPendingPerReceiver,
GraceSeconds: int64(c.GraceSeconds),
}
}
// App 实现 Service 的提交路径(M1);其余方法暂返回未实现或空操作。
type App struct {
db *store.DB
lim Limits
hash auth.HashPool
locks auth.LoginLocks
nowFn func() time.Time
rates *rateLimiter
}
// Option 配置 App。
type Option func(*App)
// WithNow 注入时钟(测试用)。
func WithNow(now func() time.Time) Option {
return func(a *App) { a.nowFn = now }
}
// WithLocks 注入对话密码锁定计数器;nil 表示不锁定。
func WithLocks(locks auth.LoginLocks) Option {
return func(a *App) { a.locks = locks }
}
// New 创建消息服务实现。hash 用于校验对话密码;locks 可为 nil。
func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App {
if lim.RequestBurst <= 0 {
lim.RequestBurst = defaultRequestBurst
}
a := &App{
db: db,
lim: lim,
hash: hash,
nowFn: time.Now,
rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst),
}
for _, opt := range opts {
opt(a)
}
return a
}
func (a *App) now() time.Time {
return a.nowFn()
}
func (a *App) protocolLimits() protocol.Limits {
return protocol.Limits{
MaxBodyBytes: a.lim.MaxBodyBytes,
MaxMetaBytes: a.lim.MaxMetaBytes,
MaxFrameBytes: a.lim.MaxFrameBytes,
}
}
func (a *App) Ack(context.Context, string, *protocol.Ack) (AckResult, error) {
return AckResult{}, ErrNotImplemented
}
func (a *App) Recall(context.Context, string, *protocol.Recall) (protocol.RecallData, error) {
return protocol.RecallData{}, ErrNotImplemented
}
func (a *App) Status(context.Context, string, *protocol.Status) (any, error) {
return nil, ErrNotImplemented
}
func (a *App) ReceiptAck(context.Context, string, *protocol.ReceiptAck) error {
return ErrNotImplemented
}
func (a *App) DispatchDue(context.Context, int64, int) (int, error) {
return 0, nil
}
func (a *App) PushPending(context.Context, string, port.ConnID) error {
return nil
}
func (a *App) OnPublishDropped(context.Context, string, port.ConnID, []byte) error {
return nil
}
func (a *App) CleanupOnce(context.Context, int64) error { return nil }
func (a *App) RecoverOnStart(context.Context) error { return nil }
func (a *App) WakePush(string) {}
var _ Service = (*App)(nil)
+52
View File
@@ -0,0 +1,52 @@
package message
import (
"database/sql"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
// dispatchMinimalTx 是 M1 最小分发:单聊插一条 pending;群按当前成员去掉发送者各插 pending;
// 消息改为 dispatched。完整 7.4 规则见 DEVIATIONS「消息 M」。
func dispatchMinimalTx(tx *sql.Tx, seq int64, senderID, destKind, destID string, sendAt int64, keep int, nowMs int64) (string, error) {
recipients := make([]string, 0, 8)
switch destKind {
case protocol.TargetEndpoint:
recipients = append(recipients, destID)
case protocol.TargetGroup:
rows, err := tx.Query(`
SELECT endpoint_id FROM group_members WHERE group_id = ? AND endpoint_id != ?`, destID, senderID)
if err != nil {
return "", err
}
defer func() { _ = rows.Close() }()
for rows.Next() {
var id string
if err := rows.Scan(&id); err != nil {
return "", err
}
recipients = append(recipients, id)
}
if err := rows.Err(); err != nil {
return "", err
}
default:
return "", errCode(protocol.CodeBadRequest, "invalid dest_kind")
}
for _, ep := range recipients {
if _, err := tx.Exec(`
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, expire_at, pushed_conn, pushed_at, attempts, updated_at)
VALUES(?,?,?,?,?,?,NULL,NULL,NULL,0,?)`,
seq, ep, sendAt, keep, DeliveryPending, "", nowMs,
); err != nil {
return "", err
}
}
state := StateDispatched
if _, err := tx.Exec(`UPDATE messages SET state = ? WHERE seq = ?`, state, seq); err != nil {
return "", err
}
return state, nil
}
+7
View File
@@ -0,0 +1,7 @@
package message
import "git.asio.asia/nixevol/NixMsg/internal/protocol"
func errCode(code, msg string) *protocol.Error {
return &protocol.Error{Code: code, Message: msg}
}
+59
View File
@@ -0,0 +1,59 @@
package message
import (
"sync"
"time"
)
// rateLimiter 是每端一个令牌桶:速率 rps、容量 burst。
// rps<=0 表示不限速。
type rateLimiter struct {
rps float64
burst float64
mu sync.Mutex
m map[string]*tokenBucket
}
type tokenBucket struct {
tokens float64
last time.Time
}
func newRateLimiter(rps float64, burst int) *rateLimiter {
if burst <= 0 {
burst = defaultRequestBurst
}
return &rateLimiter{
rps: rps,
burst: float64(burst),
m: make(map[string]*tokenBucket),
}
}
// allow 消耗 1 个令牌;允许则 true。
func (r *rateLimiter) allow(endpointID string, now time.Time) bool {
if r == nil || r.rps <= 0 {
return true
}
r.mu.Lock()
defer r.mu.Unlock()
b := r.m[endpointID]
if b == nil {
b = &tokenBucket{tokens: r.burst, last: now}
r.m[endpointID] = b
}
elapsed := now.Sub(b.last).Seconds()
if elapsed > 0 {
b.tokens += elapsed * r.rps
if b.tokens > r.burst {
b.tokens = r.burst
}
b.last = now
}
if b.tokens < 1 {
return false
}
b.tokens--
return true
}
-25
View File
@@ -1,25 +0,0 @@
package message
import (
"context"
"errors"
"testing"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
func TestStubSubmitNotImplemented(t *testing.T) {
s := NewStub()
_, err := s.Submit(context.Background(), "a", port.ConnInfo{}, &protocol.Send{})
if !errors.Is(err, ErrNotImplemented) {
t.Fatalf("got %v", err)
}
}
func TestStubRecoverNoop(t *testing.T) {
s := NewStub()
if err := s.RecoverOnStart(context.Background()); err != nil {
t.Fatal(err)
}
}
+560
View File
@@ -0,0 +1,560 @@
package message
import (
"context"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
)
// Submit 处理发送提交(DEVELOPMENT 7.3):防重 → 校验/配额/授权 → 写入 → 到点则最小分发。
func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, req *protocol.Send) (SubmitResult, error) {
if req == nil {
return SubmitResult{}, errCode(protocol.CodeBadRequest, "nil send")
}
if senderID == "" || !protocol.ValidEndpointID(senderID) {
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid sender")
}
now := a.now()
if !a.rates.allow(senderID, now) {
return SubmitResult{}, errCode(protocol.CodeRateLimited, "request rate exceeded")
}
if err := req.Validate(a.protocolLimits()); err != nil {
return SubmitResult{}, err
}
fpHex, err := protocol.RequestFingerprint(req)
if err != nil {
return SubmitResult{}, err
}
fp, err := hex.DecodeString(fpHex)
if err != nil || len(fp) != 32 {
return SubmitResult{}, fmt.Errorf("message: fingerprint decode: %w", err)
}
body, err := protocol.DecodeBody(req.Body)
if err != nil {
return SubmitResult{}, err
}
metaJSON, err := protocol.MetaCanonicalJSON(req.Meta)
if err != nil {
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid meta")
}
contentType := protocol.EffectiveContentType(req.Body)
keep := protocol.EffectiveOfflineKeep(req)
ttl := protocol.EffectiveOfflineTTL(req)
receipt := protocol.EffectiveReceipt(req)
if keep && a.lim.MaxTTLSeconds > 0 && ttl > a.lim.MaxTTLSeconds {
return SubmitResult{}, errCode(protocol.CodeBadRequest, "ttl_seconds exceeds max_ttl_seconds")
}
nowMs := now.UnixMilli()
// 防重命中可在读连接快速返回;写路径仍会再查一次以防竞态。
if res, hit, lookupErr := a.lookupIdempotent(ctx, senderID, req.ID, fp); lookupErr != nil {
return SubmitResult{}, lookupErr
} else if hit {
return res, nil
}
sender, err := a.loadEndpoint(ctx, senderID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "sender not found")
}
return SubmitResult{}, err
}
sendAt, err := a.computeSendAt(req, sender.DefaultDelayMs, nowMs)
if err != nil {
return SubmitResult{}, err
}
var (
needPassword bool
talkPHC string
targetEp *endpointRow
)
switch req.To.Kind {
case protocol.TargetEndpoint:
targetEp, err = a.loadEndpoint(ctx, req.To.ID)
if err != nil {
if errors.Is(err, sql.ErrNoRows) {
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "target not found")
}
return SubmitResult{}, err
}
if targetEp.Enabled == 0 {
return SubmitResult{}, errCode(protocol.CodeEndpointDisabled, "target disabled")
}
if senderID != req.To.ID {
needPassword, talkPHC, _, err = a.dmAuthNeeded(ctx, senderID, targetEp)
if err != nil {
return SubmitResult{}, err
}
}
case protocol.TargetGroup:
exists, member, gErr := a.groupMembership(ctx, req.To.ID, senderID)
if gErr != nil {
return SubmitResult{}, gErr
}
if !exists {
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "group not found")
}
if !member {
return SubmitResult{}, errCode(protocol.CodeNotMember, "not a group member")
}
default:
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid to.kind")
}
passwordVerified := false
if needPassword {
if locked, _ := a.talkLocked(senderID, req.To.ID, conn.RemoteIP); locked {
return SubmitResult{}, errCode(protocol.CodeRateLimited, "talk password locked")
}
if req.TalkPassword == "" {
return SubmitResult{}, errCode(protocol.CodeTalkPasswordRequired, "talk password required")
}
if a.hash == nil {
return SubmitResult{}, fmt.Errorf("message: hash pool required")
}
ok, vErr := a.hash.Verify(ctx, auth.PasswordTalk, req.TalkPassword, talkPHC)
if vErr != nil {
return SubmitResult{}, vErr
}
if !ok {
a.talkFail(senderID, req.To.ID, conn.RemoteIP)
return SubmitResult{}, errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
}
passwordVerified = true
a.talkClear(senderID, req.To.ID)
}
keepInt := 0
if keep {
keepInt = 1
}
receiptInt := 0
if receipt {
receiptInt = 1
}
var result SubmitResult
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
if res, hit, e := lookupIdempotentTx(tx, senderID, req.ID, fp); e != nil {
return e
} else if hit {
result = res
return nil
}
if e := checkQuotaTx(tx, senderID, a.lim.MaxPendingPerSender); e != nil {
return e
}
// 写事务内再确认目标与授权(防并发停用/退群)。
switch req.To.Kind {
case protocol.TargetEndpoint:
ep, e := loadEndpointTx(tx, req.To.ID)
if e != nil {
if errors.Is(e, sql.ErrNoRows) {
return errCode(protocol.CodeInvalidTarget, "target not found")
}
return e
}
if ep.Enabled == 0 {
return errCode(protocol.CodeEndpointDisabled, "target disabled")
}
if senderID != req.To.ID {
needed, phc, ver, ae := dmAuthNeededTx(tx, senderID, ep)
if ae != nil {
return ae
}
if needed {
if !passwordVerified {
if req.TalkPassword == "" {
return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
}
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
}
// 密码版本在校验后变化则拒绝,避免写过期授权。
if ep.TalkHash == nil || *ep.TalkHash != phc || ep.TalkVersion != ver {
return errCode(protocol.CodeTalkPasswordInvalid, "talk password changed")
}
if ge := upsertGrantTx(tx, senderID, req.To.ID, ver, GrantKindPassword, nowMs); ge != nil {
return ge
}
} else if passwordVerified {
// 已有授权或未设防:带对密码时仍可刷新授权(文档:带对了则写入或更新)。
if ep.TalkHash != nil && *ep.TalkHash != "" {
if ge := upsertGrantTx(tx, senderID, req.To.ID, ep.TalkVersion, GrantKindPassword, nowMs); ge != nil {
return ge
}
}
}
}
// 发送方设了对话密码且发给别人的单聊:给对方写回复授权。
snd, se := loadEndpointTx(tx, senderID)
if se != nil {
return se
}
if senderID != req.To.ID && snd.TalkHash != nil && *snd.TalkHash != "" {
if ge := upsertGrantTx(tx, req.To.ID, senderID, snd.TalkVersion, GrantKindReply, nowMs); ge != nil {
return ge
}
}
case protocol.TargetGroup:
exists, member, e := groupMembershipTx(tx, req.To.ID, senderID)
if e != nil {
return e
}
if !exists {
return errCode(protocol.CodeInvalidTarget, "group not found")
}
if !member {
return errCode(protocol.CodeNotMember, "not a group member")
}
}
state := StateScheduled
res, e := tx.Exec(`
INSERT INTO messages(
id, sender_id, dest_kind, dest_id, meta, content_type, body_enc,
send_at, keep, ttl_seconds, receipt, state, reason, created_at
) VALUES(?,?,?,?,?,?,?,?,?,?,?,?, '', ?)`,
req.ID, senderID, req.To.Kind, req.To.ID, string(metaJSON), contentType, req.Body.Enc,
sendAt, keepInt, ttl, receiptInt, state, nowMs,
)
if e != nil {
return e
}
seq, e := res.LastInsertId()
if e != nil {
return e
}
if _, e = tx.Exec(`INSERT INTO message_bodies(seq, body) VALUES(?, ?)`, seq, body); e != nil {
return e
}
if _, e = tx.Exec(
`INSERT INTO send_keys(sender_id, msg_id, request_sha256, created_at) VALUES(?,?,?,?)`,
senderID, req.ID, fp, nowMs,
); e != nil {
return e
}
finalState := state
if sendAt <= nowMs {
finalState, e = dispatchMinimalTx(tx, seq, senderID, req.To.Kind, req.To.ID, sendAt, keepInt, nowMs)
if e != nil {
return e
}
}
result = SubmitResult{ID: req.ID, SendAtMs: sendAt, State: finalState}
return nil
})
if err != nil {
return SubmitResult{}, err
}
return result, nil
}
func (a *App) computeSendAt(req *protocol.Send, defaultDelayMs, nowMs int64) (int64, error) {
if req.SendAtMs != nil && req.DelayMs != nil {
return 0, errCode(protocol.CodeBadRequest, "send_at_ms and delay_ms are mutually exclusive")
}
var sendAt int64
switch {
case req.SendAtMs != nil:
sendAt = *req.SendAtMs
case req.DelayMs != nil:
if *req.DelayMs < 0 {
return 0, errCode(protocol.CodeBadRequest, "delay_ms negative")
}
sendAt = nowMs + *req.DelayMs
default:
if defaultDelayMs < 0 {
defaultDelayMs = 0
}
sendAt = nowMs + defaultDelayMs
}
if a.lim.MaxScheduleSeconds > 0 {
maxAt := nowMs + a.lim.MaxScheduleSeconds*1000
if sendAt > maxAt {
return 0, errCode(protocol.CodeBadRequest, "send time exceeds max_schedule_seconds")
}
}
return sendAt, nil
}
type endpointRow struct {
ID string
DefaultDelayMs int64
TalkHash *string
TalkVersion int64
Enabled int
}
func (a *App) loadEndpoint(ctx context.Context, id string) (*endpointRow, error) {
row := a.db.Read.QueryRowContext(ctx, `
SELECT id, default_delay_ms, talk_hash, talk_version, enabled
FROM endpoints WHERE id = ?`, id)
return scanEndpoint(row)
}
func loadEndpointTx(tx *sql.Tx, id string) (*endpointRow, error) {
row := tx.QueryRow(`
SELECT id, default_delay_ms, talk_hash, talk_version, enabled
FROM endpoints WHERE id = ?`, id)
return scanEndpoint(row)
}
func scanEndpoint(row *sql.Row) (*endpointRow, error) {
var ep endpointRow
var talk sql.NullString
if err := row.Scan(&ep.ID, &ep.DefaultDelayMs, &talk, &ep.TalkVersion, &ep.Enabled); err != nil {
return nil, err
}
if talk.Valid {
s := talk.String
ep.TalkHash = &s
}
return &ep, nil
}
func (a *App) groupMembership(ctx context.Context, groupID, endpointID string) (exists, member bool, err error) {
var one int
err = a.db.Read.QueryRowContext(ctx, `SELECT 1 FROM groups WHERE id = ?`, groupID).Scan(&one)
if errors.Is(err, sql.ErrNoRows) {
return false, false, nil
}
if err != nil {
return false, false, err
}
err = a.db.Read.QueryRowContext(ctx,
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, groupID, endpointID,
).Scan(&one)
if errors.Is(err, sql.ErrNoRows) {
return true, false, nil
}
if err != nil {
return true, false, err
}
return true, true, nil
}
func groupMembershipTx(tx *sql.Tx, groupID, endpointID string) (exists, member bool, err error) {
var one int
err = tx.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, groupID).Scan(&one)
if errors.Is(err, sql.ErrNoRows) {
return false, false, nil
}
if err != nil {
return false, false, err
}
err = tx.QueryRow(
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, groupID, endpointID,
).Scan(&one)
if errors.Is(err, sql.ErrNoRows) {
return true, false, nil
}
if err != nil {
return true, false, err
}
return true, true, nil
}
// dmAuthNeeded 返回是否需要对话密码,以及对方当前 talk_hash / version。
func (a *App) dmAuthNeeded(ctx context.Context, senderID string, target *endpointRow) (needed bool, phc string, version int64, err error) {
if target.TalkHash == nil || *target.TalkHash == "" {
return false, "", target.TalkVersion, nil
}
ok, err := hasValidGrant(ctx, a.db.Read, senderID, target.ID, target.TalkVersion)
if err != nil {
return false, "", 0, err
}
if ok {
return false, *target.TalkHash, target.TalkVersion, nil
}
return true, *target.TalkHash, target.TalkVersion, nil
}
func dmAuthNeededTx(tx *sql.Tx, senderID string, target *endpointRow) (needed bool, phc string, version int64, err error) {
if target.TalkHash == nil || *target.TalkHash == "" {
return false, "", target.TalkVersion, nil
}
ok, err := hasValidGrantTx(tx, senderID, target.ID, target.TalkVersion)
if err != nil {
return false, "", 0, err
}
if ok {
return false, *target.TalkHash, target.TalkVersion, nil
}
return true, *target.TalkHash, target.TalkVersion, nil
}
func hasValidGrant(ctx context.Context, db *sql.DB, senderID, targetID string, talkVersion int64) (bool, error) {
var n int
err := db.QueryRowContext(ctx, `
SELECT 1 FROM talk_grants
WHERE sender_id = ? AND target_id = ? AND target_talk_version = ?
LIMIT 1`, senderID, targetID, talkVersion).Scan(&n)
if errors.Is(err, sql.ErrNoRows) {
return false, nil
}
return err == nil, err
}
func hasValidGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64) (bool, error) {
var n int
err := tx.QueryRow(`
SELECT 1 FROM talk_grants
WHERE sender_id = ? AND target_id = ? AND target_talk_version = ?
LIMIT 1`, senderID, targetID, talkVersion).Scan(&n)
if errors.Is(err, sql.ErrNoRows) {
return false, nil
}
return err == nil, err
}
func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error {
_, err := tx.Exec(`
INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at)
VALUES(?,?,?,?,?)
ON CONFLICT(sender_id, target_id) DO UPDATE SET
target_talk_version = excluded.target_talk_version,
kind = excluded.kind,
created_at = excluded.created_at`,
senderID, targetID, talkVersion, kind, nowMs,
)
return err
}
func checkQuotaTx(tx *sql.Tx, senderID string, maxPending int) error {
if maxPending <= 0 {
return nil
}
var n int
err := tx.QueryRow(`
SELECT COUNT(*) FROM messages
WHERE sender_id = ? AND state IN ('scheduled', 'dispatched')`, senderID).Scan(&n)
if err != nil {
return err
}
if n >= maxPending {
return errCode(protocol.CodeQuotaExceeded, "max_pending_per_sender exceeded")
}
return nil
}
func (a *App) lookupIdempotent(ctx context.Context, senderID, msgID string, fp []byte) (SubmitResult, bool, error) {
var stored []byte
err := a.db.Read.QueryRowContext(ctx, `
SELECT request_sha256 FROM send_keys WHERE sender_id = ? AND msg_id = ?`,
senderID, msgID,
).Scan(&stored)
if errors.Is(err, sql.ErrNoRows) {
return SubmitResult{}, false, nil
}
if err != nil {
return SubmitResult{}, false, err
}
if !bytesEqual(stored, fp) {
return SubmitResult{}, false, errCode(protocol.CodeConflict, "message id conflict")
}
res, err := loadSubmitResult(ctx, a.db.Read, senderID, msgID)
if err != nil {
return SubmitResult{}, false, err
}
return res, true, nil
}
func lookupIdempotentTx(tx *sql.Tx, senderID, msgID string, fp []byte) (SubmitResult, bool, error) {
var stored []byte
err := tx.QueryRow(`
SELECT request_sha256 FROM send_keys WHERE sender_id = ? AND msg_id = ?`,
senderID, msgID,
).Scan(&stored)
if errors.Is(err, sql.ErrNoRows) {
return SubmitResult{}, false, nil
}
if err != nil {
return SubmitResult{}, false, err
}
if !bytesEqual(stored, fp) {
return SubmitResult{}, false, errCode(protocol.CodeConflict, "message id conflict")
}
res, err := loadSubmitResultTx(tx, senderID, msgID)
if err != nil {
return SubmitResult{}, false, err
}
return res, true, nil
}
func loadSubmitResult(ctx context.Context, db *sql.DB, senderID, msgID string) (SubmitResult, error) {
var res SubmitResult
err := db.QueryRowContext(ctx, `
SELECT id, send_at, state FROM messages WHERE sender_id = ? AND id = ?`,
senderID, msgID,
).Scan(&res.ID, &res.SendAtMs, &res.State)
if errors.Is(err, sql.ErrNoRows) {
return SubmitResult{}, errCode(protocol.CodeNotFound, "idempotent key without message")
}
return res, err
}
func loadSubmitResultTx(tx *sql.Tx, senderID, msgID string) (SubmitResult, error) {
var res SubmitResult
err := tx.QueryRow(`
SELECT id, send_at, state FROM messages WHERE sender_id = ? AND id = ?`,
senderID, msgID,
).Scan(&res.ID, &res.SendAtMs, &res.State)
if errors.Is(err, sql.ErrNoRows) {
return SubmitResult{}, errCode(protocol.CodeNotFound, "idempotent key without message")
}
return res, err
}
func bytesEqual(a, b []byte) bool {
if len(a) != len(b) {
return false
}
var v byte
for i := range a {
v |= a[i] ^ b[i]
}
return v == 0
}
func (a *App) talkLocked(senderID, targetID, ip string) (bool, error) {
if a.locks == nil {
return false, nil
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: ip}); locked {
return true, nil
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
return true, nil
}
return false, nil
}
func (a *App) talkFail(senderID, targetID, ip string) {
if a.locks == nil {
return
}
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: ip})
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
}
func (a *App) talkClear(senderID, targetID string) {
if a.locks == nil {
return
}
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID})
}
+382
View File
@@ -0,0 +1,382 @@
package message
import (
"context"
"database/sql"
"errors"
"path/filepath"
"testing"
"time"
"git.asio.asia/nixevol/NixMsg/internal/app/port"
"git.asio.asia/nixevol/NixMsg/internal/auth"
"git.asio.asia/nixevol/NixMsg/internal/config"
"git.asio.asia/nixevol/NixMsg/internal/protocol"
"git.asio.asia/nixevol/NixMsg/internal/store"
)
func TestStubSubmitNotImplemented(t *testing.T) {
s := NewStub()
_, err := s.Submit(context.Background(), "a", port.ConnInfo{}, &protocol.Send{})
if !errors.Is(err, ErrNotImplemented) {
t.Fatalf("got %v", err)
}
}
func TestStubRecoverNoop(t *testing.T) {
s := NewStub()
if err := s.RecoverOnStart(context.Background()); err != nil {
t.Fatal(err)
}
}
func openTestApp(t *testing.T, lim Limits) (*App, *store.DB) {
t.Helper()
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
fixed := time.UnixMilli(1_700_000_000_000)
app := New(db, lim, auth.NewStubHashPool(),
WithNow(func() time.Time { return fixed }),
WithLocks(auth.NewStubLoginLocks()),
)
return app, db
}
func defaultTestLimits() Limits {
cfg := config.Default().Limits
lim := LimitsFromConfig(cfg)
lim.RequestsPerSecond = 0 // 测试默认不限速
return lim
}
func insertEndpoint(t *testing.T, db *store.DB, id string, talkPassword string, enabled int, defaultDelayMs int64) {
t.Helper()
ctx := context.Background()
var talk any
var talkVer int64
if talkPassword != "" {
talk = "stub$" + talkPassword
talkVer = 1
}
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(`
INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
VALUES(?,?,?,?,?,?,?,?)`,
id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, 1_700_000_000_000)
return e
})
if err != nil {
t.Fatal(err)
}
}
func baseSend(id, to string) *protocol.Send {
return &protocol.Send{
V: protocol.Version,
Type: protocol.TypeSend,
RID: "r1",
ID: id,
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: to},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hello"},
}
}
func protoCode(err error) string {
var pe *protocol.Error
if errors.As(err, &pe) {
return pe.Code
}
return ""
}
func TestSubmitTable(t *testing.T) {
t.Parallel()
t.Run("idempotent_hit", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ctx := context.Background()
req := baseSend("msg-1", "bob")
first, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
if err != nil {
t.Fatal(err)
}
if first.State != StateDispatched {
t.Fatalf("state=%s", first.State)
}
second, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
if err != nil {
t.Fatal(err)
}
if second != first {
t.Fatalf("want %+v got %+v", first, second)
}
var n int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE sender_id=? AND id=?`, "alice", "msg-1").Scan(&n); err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("messages=%d", n)
}
})
t.Run("conflict", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ctx := context.Background()
req := baseSend("msg-2", "bob")
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
t.Fatal(err)
}
other := baseSend("msg-2", "bob")
other.Body.Data = "other"
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, other)
if protoCode(err) != protocol.CodeConflict {
t.Fatalf("want conflict got %v", err)
}
})
t.Run("quota_exceeded", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
lim.MaxPendingPerSender = 1
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ctx := context.Background()
delay := int64(60_000)
req1 := baseSend("q1", "bob")
req1.DelayMs = &delay
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req1); err != nil {
t.Fatal(err)
}
req2 := baseSend("q2", "bob")
req2.DelayMs = &delay
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, req2)
if protoCode(err) != protocol.CodeQuotaExceeded {
t.Fatalf("want quota_exceeded got %v", err)
}
})
t.Run("auth_required_and_grant", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "alice-secret", 1, 0)
insertEndpoint(t, db, "bob", "secret", 1, 0)
ctx := context.Background()
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a1", "bob"))
if protoCode(err) != protocol.CodeTalkPasswordRequired {
t.Fatalf("want talk_password_required got %v", err)
}
bad := baseSend("a2", "bob")
bad.TalkPassword = "wrong"
_, err = app.Submit(ctx, "alice", port.ConnInfo{}, bad)
if protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("want talk_password_invalid got %v", err)
}
okReq := baseSend("a3", "bob")
okReq.TalkPassword = "secret"
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, okReq)
if err != nil {
t.Fatal(err)
}
if res.State != StateDispatched {
t.Fatalf("state=%s", res.State)
}
// 已有授权后不带密码也可发
if _, submitErr := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a4", "bob")); submitErr != nil {
t.Fatal(submitErr)
}
// 回复授权:bob→alice(因 alice 设了对话密码)
var kind string
err = db.Read.QueryRow(`
SELECT kind FROM talk_grants WHERE sender_id=? AND target_id=?`, "bob", "alice").Scan(&kind)
if err != nil {
t.Fatal(err)
}
if kind != GrantKindReply {
t.Fatalf("reply grant kind=%s", kind)
}
})
t.Run("self_skip_talk_password", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "secret", 1, 0)
ctx := context.Background()
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("self1", "alice")); err != nil {
t.Fatal(err)
}
})
t.Run("delay_and_send_at_mutex", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ctx := context.Background()
delay := int64(1000)
sendAt := int64(1_700_000_001_000)
req := baseSend("m-mutex", "bob")
req.DelayMs = &delay
req.SendAtMs = &sendAt
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
if protoCode(err) != protocol.CodeBadRequest {
t.Fatalf("want bad_request got %v", err)
}
})
t.Run("idempotent_before_disabled_check", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ctx := context.Background()
req := baseSend("pre-disable", "bob")
first, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
if err != nil {
t.Fatal(err)
}
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
_, e := tx.Exec(`UPDATE endpoints SET enabled = 0 WHERE id = ?`, "bob")
return e
})
if err != nil {
t.Fatal(err)
}
// 新消息应失败
_, err = app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("after-disable", "bob"))
if protoCode(err) != protocol.CodeEndpointDisabled {
t.Fatalf("want endpoint_disabled got %v", err)
}
// 原请求重试仍返回原结果
second, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
if err != nil {
t.Fatal(err)
}
if second != first {
t.Fatalf("want %+v got %+v", first, second)
}
})
t.Run("scheduled_not_dispatched", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ctx := context.Background()
delay := int64(10_000)
req := baseSend("sched-1", "bob")
req.DelayMs = &delay
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
if err != nil {
t.Fatal(err)
}
if res.State != StateScheduled {
t.Fatalf("state=%s", res.State)
}
var n int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries`).Scan(&n); err != nil {
t.Fatal(err)
}
if n != 0 {
t.Fatalf("deliveries=%d", n)
}
})
t.Run("group_dispatch_excludes_sender", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
insertEndpoint(t, db, "carol", "", 1, 0)
ctx := context.Background()
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
"g1", "g", "alice", 1_700_000_000_000); e != nil {
return e
}
for _, m := range []string{"alice", "bob", "carol"} {
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
"g1", m, 1_700_000_000_000); e != nil {
return e
}
}
return nil
})
if err != nil {
t.Fatal(err)
}
req := &protocol.Send{
V: protocol.Version,
Type: protocol.TypeSend,
RID: "r1",
ID: "gmsg-1",
To: protocol.Target{Kind: protocol.TargetGroup, ID: "g1"},
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
}
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
if err != nil {
t.Fatal(err)
}
if res.State != StateDispatched {
t.Fatalf("state=%s", res.State)
}
rows, err := db.Read.Query(`SELECT endpoint_id FROM deliveries ORDER BY endpoint_id`)
if err != nil {
t.Fatal(err)
}
defer func() { _ = rows.Close() }()
var got []string
for rows.Next() {
var id string
if err := rows.Scan(&id); err != nil {
t.Fatal(err)
}
got = append(got, id)
}
if len(got) != 2 || got[0] != "bob" || got[1] != "carol" {
t.Fatalf("recipients=%v", got)
}
})
t.Run("rate_limited", func(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
lim.RequestsPerSecond = 50
lim.RequestBurst = 2
app, db := openTestApp(t, lim)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "", 1, 0)
ctx := context.Background()
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r1", "bob")); err != nil {
t.Fatal(err)
}
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r2", "bob")); err != nil {
t.Fatal(err)
}
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r3", "bob"))
if protoCode(err) != protocol.CodeRateLimited {
t.Fatalf("want rate_limited got %v", err)
}
})
}