fix: 统一对话密码锁键为发送方加对方不含 IP

This commit is contained in:
Nixevol
2026-09-30 16:23:25 +08:00
parent b7c8b6ffd6
commit e294b5db71
17 changed files with 449 additions and 32 deletions
+47 -18
View File
@@ -32,15 +32,13 @@ SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&tal
return nil
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}); locked {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
// 没带密码不计锁定、也不因已锁返回 rate_limited(DEVELOPMENT 第 5 节:带密码的发送/进群才限流)。
if talkPassword == "" {
return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
}
if a.talkRateLimited(senderID, targetID) {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if a.hash == nil {
return errors.New("identity: hash pool required")
}
@@ -49,11 +47,11 @@ SELECT talk_hash, talk_version FROM endpoints WHERE id = ?`, targetID).Scan(&tal
return err
}
if !ok {
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP})
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
a.talkFail(senderID, targetID)
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
}
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP})
a.talkClearPair(senderID, targetID)
_ = remoteIP
nowMs := a.now().UnixMilli()
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
@@ -116,15 +114,12 @@ SELECT talk_hash, talk_version, enabled FROM endpoints WHERE id = ?`, targetID).
return nil
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}); locked {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if talkPassword == "" {
return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
}
if a.talkRateLimited(actorID, targetID) {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if a.hash == nil {
return errors.New("identity: hash pool required")
}
@@ -133,16 +128,50 @@ SELECT talk_hash, talk_version, enabled FROM endpoints WHERE id = ?`, targetID).
return err
}
if !ok {
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP})
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
a.talkFail(actorID, targetID)
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
}
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP})
a.talkClearPair(actorID, targetID)
// 进群校验成功不写入单聊授权(F15:进群密码与单聊授权分离)。
_ = talkVer
_ = remoteIP
return nil
}
func talkPairKey(senderID, targetID string) auth.LockKey {
return auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID}
}
func talkTargetKey(targetID string) auth.LockKey {
return auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}
}
func (a *App) talkRateLimited(senderID, targetID string) bool {
if a.locks == nil {
return false
}
if locked, _ := a.locks.Check(talkPairKey(senderID, targetID)); locked {
return true
}
locked, _ := a.locks.Check(talkTargetKey(targetID))
return locked
}
func (a *App) talkFail(senderID, targetID string) {
if a.locks == nil {
return
}
a.locks.Fail(talkPairKey(senderID, targetID))
a.locks.Fail(talkTargetKey(targetID))
}
func (a *App) talkClearPair(senderID, targetID string) {
if a.locks == nil {
return
}
a.locks.Clear(talkPairKey(senderID, targetID))
}
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)