561 lines
16 KiB
Go
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})
|
|
}
|