Files
NixMsg/internal/app/message/submit.go
T

561 lines
16 KiB
Go

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