diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index edf2445..1d2e76e 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -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 diff --git a/internal/app/message/app.go b/internal/app/message/app.go new file mode 100644 index 0000000..2c744e4 --- /dev/null +++ b/internal/app/message/app.go @@ -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) diff --git a/internal/app/message/dispatch_min.go b/internal/app/message/dispatch_min.go new file mode 100644 index 0000000..c118c26 --- /dev/null +++ b/internal/app/message/dispatch_min.go @@ -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 +} diff --git a/internal/app/message/errors.go b/internal/app/message/errors.go new file mode 100644 index 0000000..d18e821 --- /dev/null +++ b/internal/app/message/errors.go @@ -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} +} diff --git a/internal/app/message/rate.go b/internal/app/message/rate.go new file mode 100644 index 0000000..17555af --- /dev/null +++ b/internal/app/message/rate.go @@ -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 +} diff --git a/internal/app/message/service_test.go b/internal/app/message/service_test.go deleted file mode 100644 index d6f5b53..0000000 --- a/internal/app/message/service_test.go +++ /dev/null @@ -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) - } -} diff --git a/internal/app/message/submit.go b/internal/app/message/submit.go new file mode 100644 index 0000000..608034d --- /dev/null +++ b/internal/app/message/submit.go @@ -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}) +} diff --git a/internal/app/message/submit_test.go b/internal/app/message/submit_test.go new file mode 100644 index 0000000..d51f77b --- /dev/null +++ b/internal/app/message/submit_test.go @@ -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) + } + }) +}