diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 292b714..dcfe0db 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -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 diff --git a/internal/admin/endpoints.go b/internal/admin/endpoints.go index 342b3cd..2c8f440 100644 --- a/internal/admin/endpoints.go +++ b/internal/admin/endpoints.go @@ -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}) } diff --git a/internal/admin/endpoints_db.go b/internal/admin/endpoints_db.go index a32c16b..d6b778a 100644 --- a/internal/admin/endpoints_db.go +++ b/internal/admin/endpoints_db.go @@ -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 } diff --git a/internal/admin/endpoints_test.go b/internal/admin/endpoints_test.go index 7a39d06..9c0e005 100644 --- a/internal/admin/endpoints_test.go +++ b/internal/admin/endpoints_test.go @@ -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 { diff --git a/internal/admin/memlock.go b/internal/admin/memlock.go index bfd0ec6..612ef1a 100644 --- a/internal/admin/memlock.go +++ b/internal/admin/memlock.go @@ -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) } } diff --git a/internal/app/identity/lifecycle.go b/internal/app/identity/lifecycle.go index 7c81177..27c1d53 100644 --- a/internal/app/identity/lifecycle.go +++ b/internal/app/identity/lifecycle.go @@ -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) diff --git a/internal/app/identity/lifecycle_test.go b/internal/app/identity/lifecycle_test.go index f2d2a71..45db185 100644 --- a/internal/app/identity/lifecycle_test.go +++ b/internal/app/identity/lifecycle_test.go @@ -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 } diff --git a/internal/app/identity/register_test.go b/internal/app/identity/register_test.go index db51c04..861626a 100644 --- a/internal/app/identity/register_test.go +++ b/internal/app/identity/register_test.go @@ -92,7 +92,8 @@ func (l *registerIPLocker) Fail(key auth.LockKey) (bool, time.Duration) { return false, 0 } -func (l *registerIPLocker) ClearEndpoint(string) {} +func (l *registerIPLocker) ClearEndpoint(string) {} +func (l *registerIPLocker) ClearAllForEndpoint(string) {} func (l *registerIPLocker) Clear(key auth.LockKey) { l.mu.Lock() defer l.mu.Unlock() diff --git a/internal/app/identity/self_talk_test.go b/internal/app/identity/self_talk_test.go index 4a20087..e71618d 100644 --- a/internal/app/identity/self_talk_test.go +++ b/internal/app/identity/self_talk_test.go @@ -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) diff --git a/internal/app/identity/service.go b/internal/app/identity/service.go index 741c555..8bdc966 100644 --- a/internal/app/identity/service.go +++ b/internal/app/identity/service.go @@ -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) diff --git a/internal/app/identity/talk.go b/internal/app/identity/talk.go index cdc8eaa..ef8c11c 100644 --- a/internal/app/identity/talk.go +++ b/internal/app/identity/talk.go @@ -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) diff --git a/internal/app/message/submit.go b/internal/app/message/submit.go index 0d63252..c318617 100644 --- a/internal/app/message/submit.go +++ b/internal/app/message/submit.go @@ -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}) } diff --git a/internal/app/message/submit_test.go b/internal/app/message/submit_test.go index 38c3bc4..d726679 100644 --- a/internal/app/message/submit_test.go +++ b/internal/app/message/submit_test.go @@ -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) + } +} diff --git a/internal/auth/auth.go b/internal/auth/auth.go index 02b7a12..49da988 100644 --- a/internal/auth/auth.go +++ b/internal/auth/auth.go @@ -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) } diff --git a/internal/auth/auth_real_test.go b/internal/auth/auth_real_test.go index 6ba0b7f..432d621 100644 --- a/internal/auth/auth_real_test.go +++ b/internal/auth/auth_real_test.go @@ -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") + } +} diff --git a/internal/auth/locks.go b/internal/auth/locks.go index b2a4dd9..aa583ef 100644 --- a/internal/auth/locks.go +++ b/internal/auth/locks.go @@ -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() diff --git a/internal/auth/stub.go b/internal/auth/stub.go index 90eba90..7c5ce6a 100644 --- a/internal/auth/stub.go +++ b/internal/auth/stub.go @@ -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) {} // 编译期检查:假实现满足接口。