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