feat: 实现消息提交防重配额授权与最小分发
This commit is contained in:
@@ -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)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user