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 后再用 0004(rebase 时改号)。
- 影响:若 C-03 先合入并占用 0003,本文件 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 ## 后台接口 A
### A1 2026-09-30 ### A1 2026-09-30
+24
View File
@@ -581,6 +581,29 @@ func (h *Handler) handleEndpointTalkPassword(w http.ResponseWriter, r *http.Requ
return 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 var talkHash sql.NullString
if req.TalkPassword != "" { if req.TalkPassword != "" {
th, err := h.hash.Hash(r.Context(), auth.PasswordTalk, 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", "端不存在") httpx.WriteError(w, http.StatusNotFound, "not_found", "端不存在")
return return
} }
h.locks.Clear(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: id})
h.auditP(p, "endpoint_talk_password", id, "ok", ip) h.auditP(p, "endpoint_talk_password", id, "ok", ip)
httpx.WriteOK(w, map[string]any{"talk_password_set": talkHash.Valid}) 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 found = n > 0
return nil return nil
}) })
if err == nil && found {
h.locks.ClearAllForEndpoint(id)
}
return found, err return found, err
} }
+10
View File
@@ -269,6 +269,13 @@ func TestEndpointDisableKickAndUnlock(t *testing.T) {
t.Fatal("expected unlocked") 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", res = doReq(t, client, http.MethodPut, base+"/api/admin/endpoints/lock-1/talk-password",
`{"talk_password":"talk"}`, `{"talk_password":"talk"}`,
map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"}) map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"})
@@ -283,6 +290,9 @@ func TestEndpointDisableKickAndUnlock(t *testing.T) {
if !talkSet.TalkPasswordSet { if !talkSet.TalkPasswordSet {
t.Fatal("want talk_password_set true") 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 var ver int
if err := db.Read.QueryRow(`SELECT talk_version FROM endpoints WHERE id='lock-1'`).Scan(&ver); err != nil { 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 return false, 0
} }
// ClearEndpoint 实现 auth.LoginLocks。 // ClearEndpoint 实现 auth.LoginLocks:只清登录锁定。
func (l *MemoryLoginLocks) ClearEndpoint(endpointID string) { func (l *MemoryLoginLocks) ClearEndpoint(endpointID string) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
for k := range l.entries { for k := range l.entries {
// kind|endpoint|peer|ip
parts := splitLockKey(k) 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) delete(l.entries, k)
} }
} }
+3
View File
@@ -126,6 +126,9 @@ WHERE id = ?`, endpointID); e != nil {
if err != nil { if err != nil {
return err return err
} }
if hardDelete {
a.locks.ClearAllForEndpoint(endpointID)
}
a.publishRevokes(ctx, revokes) a.publishRevokes(ctx, revokes)
a.publishGroupEvents(ctx, notifies) 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 { type lifecycleKick struct {
Calls []string 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) ClearEndpoint(string) {}
func (l *registerIPLocker) ClearAllForEndpoint(string) {}
func (l *registerIPLocker) Clear(key auth.LockKey) { func (l *registerIPLocker) Clear(key auth.LockKey) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
+102
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"errors" "errors"
"fmt"
"path/filepath" "path/filepath"
"testing" "testing"
"time" "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) { func TestUnlockSelfAndNoPassword(t *testing.T) {
t.Parallel() t.Parallel()
app, db, _ := openIdentity(t) 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) SelfChangeLoginPassword(ctx context.Context, endpointID, oldPassword, newPassword, remoteIP string) (sessionToken string, err error)
SelfLogout(ctx context.Context, endpointID string) 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 UnlockTalk(ctx context.Context, senderID, targetID, talkPassword, remoteIP string) error
// HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。 // HasTalkGrant 查询发送方对目标是否有有效授权(无对话密码或已有匹配版本授权)。
HasTalkGrant(ctx context.Context, senderID, targetID string) (bool, error) 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 return nil
} }
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}); locked { // 没带密码不计锁定、也不因已锁返回 rate_limited(DEVELOPMENT 第 5 节:带密码的发送/进群才限流)。
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 == "" { if talkPassword == "" {
return errCode(protocol.CodeTalkPasswordRequired, "talk password required") return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
} }
if a.talkRateLimited(senderID, targetID) {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if a.hash == nil { if a.hash == nil {
return errors.New("identity: hash pool required") 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 return err
} }
if !ok { if !ok {
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: remoteIP}) a.talkFail(senderID, targetID)
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") 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() nowMs := a.now().UnixMilli()
return a.db.Queue.Do(ctx, func(tx *sql.Tx) error { 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 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 == "" { if talkPassword == "" {
return errCode(protocol.CodeTalkPasswordRequired, "talk password required") return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
} }
if a.talkRateLimited(actorID, targetID) {
return errCode(protocol.CodeRateLimited, "talk password locked")
}
if a.hash == nil { if a.hash == nil {
return errors.New("identity: hash pool required") return errors.New("identity: hash pool required")
} }
@@ -133,16 +128,50 @@ SELECT talk_hash, talk_version, enabled FROM endpoints WHERE id = ?`, targetID).
return err return err
} }
if !ok { if !ok {
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: actorID, PeerID: targetID, IP: remoteIP}) a.talkFail(actorID, targetID)
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") 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:进群密码与单聊授权分离)。 // 进群校验成功不写入单聊授权(F15:进群密码与单聊授权分离)。
_ = talkVer _ = talkVer
_ = remoteIP
return nil 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 { func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error {
_, err := tx.Exec(` _, err := tx.Exec(`
INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at) 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 passwordVerified := false
if needPassword { if needPassword {
if locked, _ := a.talkLocked(senderID, req.To.ID, conn.RemoteIP); locked {
return SubmitResult{}, errCode(protocol.CodeRateLimited, "talk password locked")
}
if req.TalkPassword == "" { if req.TalkPassword == "" {
return SubmitResult{}, errCode(protocol.CodeTalkPasswordRequired, "talk password required") 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 { if a.hash == nil {
return SubmitResult{}, fmt.Errorf("message: hash pool required") 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 return SubmitResult{}, vErr
} }
if !ok { if !ok {
a.talkFail(senderID, req.To.ID, conn.RemoteIP) a.talkFail(senderID, req.To.ID)
return SubmitResult{}, errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid") return SubmitResult{}, errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
} }
passwordVerified = true passwordVerified = true
@@ -572,11 +572,11 @@ func bytesEqual(a, b []byte) bool {
return v == 0 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 { if a.locks == nil {
return false, 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 return true, nil
} }
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked { 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 return false, nil
} }
func (a *App) talkFail(senderID, targetID, ip string) { func (a *App) talkFail(senderID, targetID string) {
if a.locks == nil { if a.locks == nil {
return 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}) a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
} }
+42
View File
@@ -4,6 +4,7 @@ import (
"context" "context"
"database/sql" "database/sql"
"errors" "errors"
"fmt"
"math" "math"
"path/filepath" "path/filepath"
"testing" "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" LockLoginEndpointIP LockKind = "login_endpoint_ip"
// LockLoginEndpoint:编号总数,1 小时内 50 次错 → 暂停该编号密码登录 1 小时。 // LockLoginEndpoint:编号总数,1 小时内 50 次错 → 暂停该编号密码登录 1 小时。
LockLoginEndpoint LockKind = "login_endpoint" LockLoginEndpoint LockKind = "login_endpoint"
// LockTalkPair:发送方 + 对方对话密码。 // LockTalkPair:发送方 + 对方对话密码(不含 IP)。
LockTalkPair LockKind = "talk_pair" LockTalkPair LockKind = "talk_pair"
// LockTalkTarget:对方对话密码总数。 // LockTalkTarget:对方对话密码总数。
LockTalkTarget LockKind = "talk_target" LockTalkTarget LockKind = "talk_target"
@@ -90,6 +90,8 @@ type LoginLocks interface {
Fail(key LockKey) (locked bool, retryAfter time.Duration) Fail(key LockKey) (locked bool, retryAfter time.Duration)
// ClearEndpoint 清除某端编号相关的登录锁定(两种都清),对应管理 unlock。 // ClearEndpoint 清除某端编号相关的登录锁定(两种都清),对应管理 unlock。
ClearEndpoint(endpointID string) ClearEndpoint(endpointID string)
// ClearAllForEndpoint 清除该编号作为 EndpointID 或 PeerID 出现的全部锁定(删除端后防编号复用继承)。
ClearAllForEndpoint(endpointID string)
// Clear 清除精确键。 // Clear 清除精确键。
Clear(key LockKey) Clear(key LockKey)
} }
+51
View File
@@ -138,3 +138,54 @@ func TestLoginLocksClearEndpoint(t *testing.T) {
t.Fatal("e1 total lock should clear") 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 清除精确键。 // Clear 清除精确键。
func (l *MemoryLocks) Clear(key LockKey) { func (l *MemoryLocks) Clear(key LockKey) {
l.mu.Lock() l.mu.Lock()
+6
View File
@@ -94,6 +94,12 @@ func (l *StubLoginLocks) ClearEndpoint(endpointID string) {
l.Cleared = append(l.Cleared, endpointID) 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) {} func (l *StubLoginLocks) Clear(LockKey) {}
// 编译期检查:假实现满足接口。 // 编译期检查:假实现满足接口。