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 { var claimed bool finalState, claimed, e = a.dispatchFullTx(tx, seq, senderID, req.To.Kind, req.To.ID, sendAt, keepInt, ttl, receipt, nowMs) if e != nil { return e } _ = claimed } result = SubmitResult{ID: req.ID, SendAtMs: sendAt, State: finalState} return nil }) if err != nil { return SubmitResult{}, err } if result.State == StateDispatched { a.wakeReceivers(ctx, result.ID, senderID) } return result, nil } func (a *App) wakeReceivers(ctx context.Context, msgID, senderID string) { rows, err := a.db.Read.QueryContext(ctx, ` SELECT d.endpoint_id FROM deliveries d JOIN messages m ON m.seq = d.seq WHERE m.sender_id = ? AND m.id = ? AND d.state = 'pending'`, senderID, msgID) if err != nil { return } defer func() { _ = rows.Close() }() for rows.Next() { var ep string if rows.Scan(&ep) == nil { a.WakePush(ep) } } } 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}) }