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
+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
}
+2 -1
View File
@@ -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()
+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)
}
}