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
+9
View File
@@ -778,6 +778,15 @@
- 备选方案:等 C-03 占用 0003 后再用 0004(rebase 时改号)。
- 影响:若 C-03 先合入并占用 0003,本文件 rebase 时改号。
### 复审修复 U-03
1. **对话密码锁键统一为发送方+对方,不含 IP**
- 原条款:DEVELOPMENT 第 5 节按「发送方 + 对方」计数;锁定期间 unlock、**带密码的**发送和拉人进群返回 `rate_limited`。PRD F15 / D24:改密清按对方计的总数;删除后编号可复用。issue #41。
- 实际做法:`LockTalkPair` 键为 `{Kind, 发送方, 对方}`,IP 留空。unlock / send / 进群共用该键;没带密码直接 `talk_password_required`,不计次、不因已锁改成 `rate_limited`。`UnlockTalk`/`CheckTalkPasswordForJoin` 仍保留 `remoteIP` 参数以免改接线签名。后台 `PUT talk-password` 有 Identity 时调 `SelfSetTalkPassword`(已清 `LockTalkTarget`),无 Identity 时改库后 `Clear(LockTalkTarget)`。删除端调用 `LoginLocks.ClearAllForEndpoint`。不改 group `emit`,不修 H-02。
- 原因:原先 unlock 带 IP、send 成功清零不带 IP、群加人用空 IP,同一发送方换 IP 可再猜 10 次;没带密码也会被已锁挡成 `rate_limited`,SDK 会退避最多 1 小时。
- 备选方案:对话密码也按编号+IP(否决,与第 5 节原文及 F15 不一致)。
- 影响:同一对端合计 10 次错即锁;没带密码始终是 `talk_password_required`;后台改密可解除按对方暂停;删端后同编号重开不继承锁定。
## 后台接口 A
### A1 2026-09-30
+24
View File
@@ -581,6 +581,29 @@ func (h *Handler) handleEndpointTalkPassword(w http.ResponseWriter, r *http.Requ
return
}
if h.identity != nil {
err := h.identity.SelfSetTalkPassword(r.Context(), id, req.TalkPassword)
if err != nil {
if isEndpointNotFound(err) {
h.auditP(p, "endpoint_talk_password", id, "not_found", ip)
httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
return
}
var pe *protocol.Error
if errors.As(err, &pe) && pe.Code == protocol.CodeBadRequest {
h.auditP(p, "endpoint_talk_password", id, "bad_request", ip)
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "对话密码不合法")
return
}
h.auditP(p, "endpoint_talk_password", id, "error", ip)
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
return
}
h.auditP(p, "endpoint_talk_password", id, "ok", ip)
httpx.WriteOK(w, map[string]any{"talk_password_set": req.TalkPassword != ""})
return
}
var talkHash sql.NullString
if req.TalkPassword != "" {
th, err := h.hash.Hash(r.Context(), auth.PasswordTalk, req.TalkPassword)
@@ -602,6 +625,7 @@ func (h *Handler) handleEndpointTalkPassword(w http.ResponseWriter, r *http.Requ
httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
return
}
h.locks.Clear(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: id})
h.auditP(p, "endpoint_talk_password", id, "ok", ip)
httpx.WriteOK(w, map[string]any{"talk_password_set": talkHash.Valid})
}
+3
View File
@@ -370,6 +370,9 @@ func (h *Handler) deleteEndpointBasic(ctx context.Context, id string) (found boo
found = n > 0
return nil
})
if err == nil && found {
h.locks.ClearAllForEndpoint(id)
}
return found, err
}
+10
View File
@@ -269,6 +269,13 @@ func TestEndpointDisableKickAndUnlock(t *testing.T) {
t.Fatal("expected unlocked")
}
for i := 0; i < 50; i++ {
locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "lock-1"})
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "lock-1"}); !locked {
t.Fatal("expected talk target lock before talk-password")
}
res = doReq(t, client, http.MethodPut, base+"/api/admin/endpoints/lock-1/talk-password",
`{"talk_password":"talk"}`,
map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"})
@@ -283,6 +290,9 @@ func TestEndpointDisableKickAndUnlock(t *testing.T) {
if !talkSet.TalkPasswordSet {
t.Fatal("want talk_password_set true")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "lock-1"}); locked {
t.Fatal("talk-password should clear LockTalkTarget")
}
var ver int
if err := db.Read.QueryRow(`SELECT talk_version FROM endpoints WHERE id='lock-1'`).Scan(&ver); err != nil {
+27 -3
View File
@@ -89,14 +89,38 @@ func (l *MemoryLoginLocks) Fail(key auth.LockKey) (bool, time.Duration) {
return false, 0
}
// ClearEndpoint 实现 auth.LoginLocks。
// ClearEndpoint 实现 auth.LoginLocks:只清登录锁定。
func (l *MemoryLoginLocks) ClearEndpoint(endpointID string) {
l.mu.Lock()
defer l.mu.Unlock()
for k := range l.entries {
// kind|endpoint|peer|ip
parts := splitLockKey(k)
if len(parts) >= 2 && parts[1] == endpointID {
if len(parts) != 4 {
continue
}
kind, ep := auth.LockKind(parts[0]), parts[1]
if ep != endpointID {
continue
}
if kind == auth.LockLoginEndpointIP || kind == auth.LockLoginEndpoint {
delete(l.entries, k)
}
}
}
// ClearAllForEndpoint 实现 auth.LoginLocks。
func (l *MemoryLoginLocks) ClearAllForEndpoint(endpointID string) {
l.mu.Lock()
defer l.mu.Unlock()
if endpointID == "" {
return
}
for k := range l.entries {
parts := splitLockKey(k)
if len(parts) != 4 {
continue
}
if parts[1] == endpointID || parts[2] == endpointID {
delete(l.entries, k)
}
}
+3
View File
@@ -126,6 +126,9 @@ WHERE id = ?`, endpointID); e != nil {
if err != nil {
return err
}
if hardDelete {
a.locks.ClearAllForEndpoint(endpointID)
}
a.publishRevokes(ctx, revokes)
a.publishGroupEvents(ctx, notifies)
+92
View File
@@ -501,6 +501,98 @@ func TestAdminDisableDeleteHTTP(t *testing.T) {
}
}
func TestU03AdminTalkPasswordHTTP(t *testing.T) {
t.Parallel()
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
fixed := time.UnixMilli(1_700_000_000_000)
hash := auth.NewStubHashPool()
if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, "adminpassword1"); seedErr != nil {
t.Fatal(seedErr)
}
locks := auth.NewLoginLocks()
locks.SetClock(func() time.Time { return fixed })
idApp := identity.New(identity.Config{
DB: db, Hash: hash, Locks: locks, Sessions: auth.NewSessionTokens(),
MaxScheduleSeconds: 86400, Now: func() time.Time { return fixed },
})
h := admin.New(admin.Deps{
DB: db, Hash: hash, Tokens: admin.NewRandomAPITokens(),
Locks: locks, Identity: idApp,
})
srv := httptest.NewServer(h)
t.Cleanup(srv.Close)
jar, _ := cookiejar.New(nil)
client := &http.Client{Jar: jar}
loginRes, err := client.Post(srv.URL+"/api/admin/login", "application/json",
strings.NewReader(`{"username":"admin","password":"adminpassword1"}`))
if err != nil {
t.Fatal(err)
}
_ = loginRes.Body.Close()
if loginRes.StatusCode != 200 {
t.Fatalf("login %d", loginRes.StatusCode)
}
req, _ := http.NewRequest(http.MethodPost, srv.URL+"/api/admin/endpoints",
strings.NewReader(`{"id":"alice","name":"alice","login_password":"password12"}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Nixmsg-Request", "1")
res, err := client.Do(req)
if err != nil {
t.Fatal(err)
}
raw, _ := io.ReadAll(res.Body)
_ = res.Body.Close()
if res.StatusCode != 200 {
t.Fatalf("create: %d %s", res.StatusCode, raw)
}
for i := 0; i < 50; i++ {
locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"})
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"}); !locked {
t.Fatal("expected target lock")
}
req, _ = http.NewRequest(http.MethodPut, srv.URL+"/api/admin/endpoints/alice/talk-password",
strings.NewReader(`{"talk_password":"secret"}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Nixmsg-Request", "1")
res, err = client.Do(req)
if err != nil {
t.Fatal(err)
}
raw, _ = io.ReadAll(res.Body)
_ = res.Body.Close()
if res.StatusCode != 200 {
t.Fatalf("talk-password: %d %s", res.StatusCode, raw)
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "alice"}); locked {
t.Fatal("identity SelfSetTalkPassword should clear LockTalkTarget")
}
req, _ = http.NewRequest(http.MethodPut, srv.URL+"/api/admin/endpoints/missing/talk-password",
strings.NewReader(`{"talk_password":"secret"}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Nixmsg-Request", "1")
res, err = client.Do(req)
if err != nil {
t.Fatal(err)
}
raw, _ = io.ReadAll(res.Body)
_ = res.Body.Close()
if res.StatusCode != http.StatusNotFound {
t.Fatalf("missing endpoint want 404 got %d %s", res.StatusCode, raw)
}
}
type lifecycleKick struct {
Calls []string
}
+1
View File
@@ -93,6 +93,7 @@ func (l *registerIPLocker) Fail(key auth.LockKey) (bool, time.Duration) {
}
func (l *registerIPLocker) ClearEndpoint(string) {}
func (l *registerIPLocker) ClearAllForEndpoint(string) {}
func (l *registerIPLocker) Clear(key auth.LockKey) {
l.mu.Lock()
defer l.mu.Unlock()
+102
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"errors"
"fmt"
"path/filepath"
"testing"
"time"
@@ -297,6 +298,107 @@ func TestSelfLoginPasswordAndLogout(t *testing.T) {
}
}
func TestU03TalkPairLockIgnoresIPAndEmptyPassword(t *testing.T) {
t.Parallel()
app, db, locks := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
for i := 0; i < 9; i++ {
err := app.UnlockTalk(ctx, "alice", "bob", "wrong", fmt.Sprintf("10.0.0.%d", i+1))
if protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("fail %d: %v", i, err)
}
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "9.9.9.9"); err != nil {
t.Fatalf("9th+correct from other IP should succeed: %v", err)
}
for i := 0; i < 5; i++ {
if err := app.UnlockTalk(ctx, "alice", "bob", "wrong", fmt.Sprintf("1.1.1.%d", i+1)); protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("unlock fail %d: %v", i, err)
}
}
for i := 0; i < 5; i++ {
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "wrong", fmt.Sprintf("2.2.2.%d", i+1)); protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("join fail %d: %v", i, err)
}
}
pair := auth.LockKey{Kind: auth.LockTalkPair, EndpointID: "alice", PeerID: "bob"}
if locked, _ := locks.Check(pair); !locked {
t.Fatal("pair should lock after 10 wrong attempts across IPs and paths")
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "8.8.8.8"); protoCode(err) != protocol.CodeRateLimited {
t.Fatalf("want rate_limited unlock got %v", err)
}
if err := app.CheckTalkPasswordForJoin(ctx, "alice", "bob", "secret", "7.7.7.7"); protoCode(err) != protocol.CodeRateLimited {
t.Fatalf("want rate_limited join got %v", err)
}
if err := app.UnlockTalk(ctx, "alice", "bob", "", "6.6.6.6"); protoCode(err) != protocol.CodeTalkPasswordRequired {
t.Fatalf("empty while locked want required got %v", err)
}
}
func TestU03AdminChangeAndDeleteClearTalkLocks(t *testing.T) {
t.Parallel()
app, db, locks := openIdentity(t)
ctx := context.Background()
insertEP(t, db, "alice", "password1")
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
for i := 0; i < 50; i++ {
id := fmt.Sprintf("u%02d", i)
insertEP(t, db, id, "password1")
_ = app.UnlockTalk(ctx, id, "bob", "wrong", "2.2.2.2")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"}); !locked {
t.Fatal("expected talk target lock")
}
if err := app.SelfSetTalkPassword(ctx, "bob", "secret2"); err != nil {
t.Fatal(err)
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"}); locked {
t.Fatal("admin/self change should clear LockTalkTarget")
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret2", "3.3.3.3"); err != nil {
t.Fatalf("unlock after change: %v", err)
}
for i := 0; i < 10; i++ {
_ = app.UnlockTalk(ctx, "alice", "bob", "wrong", fmt.Sprintf("4.4.4.%d", i+1))
}
locks.Fail(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "bob"})
locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: "bob", PeerID: "alice"})
if err := app.Delete(ctx, "bob"); err != nil {
t.Fatal(err)
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: "alice", PeerID: "bob"}); locked {
t.Fatal("delete should clear pair where bob is peer")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: "bob", PeerID: "alice"}); locked {
t.Fatal("delete should clear pair where bob is sender")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: "bob"}); locked {
t.Fatal("delete should clear LockTalkTarget")
}
if locked, _ := locks.Check(auth.LockKey{Kind: auth.LockLoginEndpoint, EndpointID: "bob"}); locked {
t.Fatal("delete should clear login lock")
}
insertEP(t, db, "bob", "password1")
if err := app.SelfSetTalkPassword(ctx, "bob", "secret"); err != nil {
t.Fatal(err)
}
if err := app.UnlockTalk(ctx, "alice", "bob", "secret", "5.5.5.5"); err != nil {
t.Fatalf("reopened bob must not inherit talk lock: %v", err)
}
}
func TestUnlockSelfAndNoPassword(t *testing.T) {
t.Parallel()
app, db, _ := openIdentity(t)
+2 -1
View File
@@ -50,7 +50,8 @@ type Service interface {
SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error)
SelfLogout(ctx context.Context, endpointID string) error
// UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock);remoteIP 计入对话密码锁定。
// UnlockTalk 校验并写入 password 类对话授权(第 6.6 节 unlock)。
// remoteIP 保留给接线方;对话密码锁键为 {LockTalkPair, 发送方, 对方},不含 IP。
UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error
// HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。
HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error)
+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)
+8 -8
View File
@@ -118,12 +118,12 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
}
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 locked, _ := a.talkLocked(senderID, req.To.ID); locked {
return SubmitResult{}, errCode(protocol.CodeRateLimited, "talk password locked")
}
if a.hash == nil {
return SubmitResult{}, fmt.Errorf("message: hash pool required")
}
@@ -132,7 +132,7 @@ func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, r
return SubmitResult{}, vErr
}
if !ok {
a.talkFail(senderID, req.To.ID, conn.RemoteIP)
a.talkFail(senderID, req.To.ID)
return SubmitResult{}, errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
}
passwordVerified = true
@@ -572,11 +572,11 @@ func bytesEqual(a, b []byte) bool {
return v == 0
}
func (a *App) talkLocked(senderID, targetID, ip string) (bool, error) {
func (a *App) talkLocked(senderID, targetID 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 {
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID}); locked {
return true, nil
}
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
@@ -585,11 +585,11 @@ func (a *App) talkLocked(senderID, targetID, ip string) (bool, error) {
return false, nil
}
func (a *App) talkFail(senderID, targetID, ip string) {
func (a *App) talkFail(senderID, targetID 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.LockTalkPair, EndpointID: senderID, PeerID: targetID})
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
}
+42
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"errors"
"fmt"
"math"
"path/filepath"
"testing"
@@ -519,3 +520,44 @@ WHERE m.id='late-1' AND d.endpoint_id='bob'`).Scan(&bobN); err != nil {
}
})
}
func TestU03SubmitTalkLockNoIP(t *testing.T) {
t.Parallel()
lim := defaultTestLimits()
dir := t.TempDir()
db, err := store.Open(filepath.Join(dir, "data"), "FULL")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
fixed := time.UnixMilli(1_700_000_000_000)
locks := auth.NewLoginLocks()
locks.SetClock(func() time.Time { return fixed })
app := New(db, lim, auth.NewStubHashPool(),
WithNow(func() time.Time { return fixed }),
WithLocks(locks),
)
insertEndpoint(t, db, "alice", "", 1, 0)
insertEndpoint(t, db, "bob", "secret", 1, 0)
ctx := context.Background()
for i := 0; i < 10; i++ {
bad := baseSend(fmt.Sprintf("w%d", i), "bob")
bad.TalkPassword = "wrong"
_, err := app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: fmt.Sprintf("10.0.0.%d", i+1)}, bad)
if protoCode(err) != protocol.CodeTalkPasswordInvalid {
t.Fatalf("i=%d got %v", i, err)
}
}
empty := baseSend("empty", "bob")
_, err = app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "8.8.8.8"}, empty)
if protoCode(err) != protocol.CodeTalkPasswordRequired {
t.Fatalf("empty while locked want required got %v", err)
}
okReq := baseSend("ok1", "bob")
okReq.TalkPassword = "secret"
_, err = app.Submit(ctx, "alice", port.ConnInfo{RemoteIP: "9.9.9.9"}, okReq)
if protoCode(err) != protocol.CodeRateLimited {
t.Fatalf("correct password while locked want rate_limited got %v", err)
}
}
+3 -1
View File
@@ -64,7 +64,7 @@ const (
LockLoginEndpointIP LockKind = "login_endpoint_ip"
// LockLoginEndpoint:编号总数,1 小时内 50 次错 → 暂停该编号密码登录 1 小时。
LockLoginEndpoint LockKind = "login_endpoint"
// LockTalkPair:发送方 + 对方对话密码。
// LockTalkPair:发送方 + 对方对话密码(不含 IP)。
LockTalkPair LockKind = "talk_pair"
// LockTalkTarget:对方对话密码总数。
LockTalkTarget LockKind = "talk_target"
@@ -90,6 +90,8 @@ type LoginLocks interface {
Fail(key LockKey) (locked bool, retryAfter time.Duration)
// ClearEndpoint 清除某端编号相关的登录锁定(两种都清),对应管理 unlock。
ClearEndpoint(endpointID string)
// ClearAllForEndpoint 清除该编号作为 EndpointID 或 PeerID 出现的全部锁定(删除端后防编号复用继承)。
ClearAllForEndpoint(endpointID string)
// Clear 清除精确键。
Clear(key LockKey)
}
+51
View File
@@ -138,3 +138,54 @@ func TestLoginLocksClearEndpoint(t *testing.T) {
t.Fatal("e1 total lock should clear")
}
}
func TestLoginLocksClearEndpointKeepsTalk(t *testing.T) {
locks := NewLoginLocks()
talk := LockKey{Kind: LockTalkPair, EndpointID: "e1", PeerID: "e2"}
target := LockKey{Kind: LockTalkTarget, EndpointID: "e1"}
login := LockKey{Kind: LockLoginEndpoint, EndpointID: "e1"}
for i := 0; i < 10; i++ {
locks.Fail(talk)
locks.Fail(login)
}
for i := 0; i < 50; i++ {
locks.Fail(target)
}
locks.ClearEndpoint("e1")
if locked, _ := locks.Check(login); locked {
t.Fatal("login lock should clear")
}
if locked, _ := locks.Check(talk); !locked {
t.Fatal("talk pair lock must survive ClearEndpoint")
}
if locked, _ := locks.Check(target); !locked {
t.Fatal("talk target lock must survive ClearEndpoint")
}
}
func TestLoginLocksClearAllForEndpoint(t *testing.T) {
locks := NewLoginLocks()
asSender := LockKey{Kind: LockTalkPair, EndpointID: "gone", PeerID: "peer"}
asPeer := LockKey{Kind: LockTalkPair, EndpointID: "other", PeerID: "gone"}
target := LockKey{Kind: LockTalkTarget, EndpointID: "gone"}
login := LockKey{Kind: LockLoginEndpointIP, EndpointID: "gone", IP: "1.1.1.1"}
keep := LockKey{Kind: LockTalkPair, EndpointID: "keep", PeerID: "peer"}
for i := 0; i < 10; i++ {
locks.Fail(asSender)
locks.Fail(asPeer)
locks.Fail(login)
locks.Fail(keep)
}
for i := 0; i < 50; i++ {
locks.Fail(target)
}
locks.ClearAllForEndpoint("gone")
for _, k := range []LockKey{asSender, asPeer, target, login} {
if locked, _ := locks.Check(k); locked {
t.Fatalf("expected %s cleared", lockMapKey(k))
}
}
if locked, _ := locks.Check(keep); !locked {
t.Fatal("unrelated pair should remain")
}
}
+18
View File
@@ -245,6 +245,24 @@ func (l *MemoryLocks) ClearEndpoint(endpointID string) {
}
}
// ClearAllForEndpoint 清除该编号作为 EndpointID 或 PeerID 出现的全部锁定。
func (l *MemoryLocks) ClearAllForEndpoint(endpointID string) {
l.mu.Lock()
defer l.mu.Unlock()
if endpointID == "" {
return
}
for k := range l.entries {
parts := splitLockKey(k)
if len(parts) != 4 {
continue
}
if parts[1] == endpointID || parts[2] == endpointID {
delete(l.entries, k)
}
}
}
// Clear 清除精确键。
func (l *MemoryLocks) Clear(key LockKey) {
l.mu.Lock()
+6
View File
@@ -94,6 +94,12 @@ func (l *StubLoginLocks) ClearEndpoint(endpointID string) {
l.Cleared = append(l.Cleared, endpointID)
}
func (l *StubLoginLocks) ClearAllForEndpoint(endpointID string) {
l.mu.Lock()
defer l.mu.Unlock()
l.Cleared = append(l.Cleared, endpointID)
}
func (l *StubLoginLocks) Clear(LockKey) {}
// 编译期检查:假实现满足接口。