fix: 统一对话密码锁键为发送方加对方不含 IP
This commit is contained in:
@@ -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})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user