merge: identity u01-u05

This commit is contained in:
Nixevol
2026-09-30 15:28:06 +08:00
14 changed files with 494 additions and 19 deletions
+60
View File
@@ -134,6 +134,66 @@ func TestRegistrationToggleAndGenerate(t *testing.T) {
}
}
func TestRegistrationEnableRequiresCode(t *testing.T) {
db, srv, client := setupA3(t)
base := srv.URL
res := doReq(t, client, http.MethodPut, base+"/api/admin/registration",
`{"enabled":true}`, csrf())
env := decodeEnv(t, res)
if res.StatusCode != 400 || env.OK {
t.Fatalf("enable without code: %d %+v", res.StatusCode, env)
}
if env.Error == nil || !strings.Contains(env.Error.Message, "8–64") {
t.Fatalf("want 8–64 message, got %+v", env.Error)
}
res = doReq(t, client, http.MethodGet, base+"/api/admin/registration", "", nil)
env = decodeEnv(t, res)
var got map[string]any
_ = json.Unmarshal(env.Data, &got)
if got["enabled"] != false {
t.Fatalf("switch must stay off after 400: %v", got)
}
var n int
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM settings WHERE key = ? AND value = '1'`, "registration_enabled").Scan(&n); err != nil {
t.Fatal(err)
}
if n != 0 {
t.Fatalf("enabled setting rolled back, count=%d", n)
}
res = doReq(t, client, http.MethodPut, base+"/api/admin/registration",
`{"enabled":true,"generate":true}`, csrf())
env = decodeEnv(t, res)
if res.StatusCode != 200 || !env.OK {
t.Fatalf("enable+generate: %d %+v", res.StatusCode, env)
}
_ = json.Unmarshal(env.Data, &got)
code, _ := got["code"].(string)
if got["enabled"] != true || len(code) != 16 {
t.Fatalf("after generate: %v", got)
}
res = doReq(t, client, http.MethodPut, base+"/api/admin/registration",
`{"enabled":false}`, csrf())
if res.StatusCode != 200 {
t.Fatalf("disable: %d", res.StatusCode)
}
_ = res.Body.Close()
res = doReq(t, client, http.MethodPut, base+"/api/admin/registration",
`{"enabled":true,"code":"abcdefgh"}`, csrf())
env = decodeEnv(t, res)
if res.StatusCode != 200 || !env.OK {
t.Fatalf("enable+code: %d %+v", res.StatusCode, env)
}
_ = json.Unmarshal(env.Data, &got)
if got["enabled"] != true || got["code"] != "abcdefgh" {
t.Fatalf("after enable+code: %v", got)
}
}
func TestMessageDetailHasNoBody(t *testing.T) {
db, srv, client := setupA3(t)
base := srv.URL
+13 -3
View File
@@ -77,9 +77,7 @@ func (h *Handler) handleRegistrationPut(w http.ResponseWriter, r *http.Request)
); e != nil {
return e
}
return nil
}
if req.Code != nil {
} else if req.Code != nil {
code := *req.Code
n := utf8.RuneCountInString(code)
if n < minRegistrationCodeLen || n > maxRegistrationCodeLen {
@@ -93,6 +91,18 @@ func (h *Handler) handleRegistrationPut(w http.ResponseWriter, r *http.Request)
return e
}
}
var enabledVal, codeVal sql.NullString
if e := tx.QueryRow(`SELECT value FROM settings WHERE key = ?`, settingRegistrationEnabled).Scan(&enabledVal); e != nil && !errors.Is(e, sql.ErrNoRows) {
return e
}
if e := tx.QueryRow(`SELECT value FROM settings WHERE key = ?`, settingRegistrationCode).Scan(&codeVal); e != nil && !errors.Is(e, sql.ErrNoRows) {
return e
}
n := utf8.RuneCountInString(codeVal.String)
if registrationTruthy(enabledVal.String) && (n < minRegistrationCodeLen || n > maxRegistrationCodeLen) {
return errBadRequest("开启自助注册前须先设置 8–64 字符的安全码")
}
return nil
})
if err != nil {
+11
View File
@@ -25,6 +25,9 @@ const (
settingRegistrationEnabled = "registration_enabled"
settingRegistrationCode = "registration_code"
minRegistrationCodeLen = 8
maxRegistrationCodeLen = 64
sourceSelf = "self"
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
@@ -144,6 +147,11 @@ func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRe
if !enabled {
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationClosed, "registration closed")
}
// 存储码不是 8–64 字符时视为关闭(含开启+空码的旧库)。放在锁定检查之前,不计入锁定。
if n := utf8.RuneCountInString(storedCode); n < minRegistrationCodeLen || n > maxRegistrationCodeLen {
h.cfg.Logger.Warn("register", "result", protocol.CodeRegistrationClosed, "reason", "unusable_code", "ip", ip)
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationClosed, "registration closed")
}
lockKey := auth.LockKey{Kind: auth.LockRegisterIP, IP: ip}
if locked, _ := h.cfg.Locks.Check(lockKey); locked {
@@ -163,6 +171,9 @@ func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRe
}
return RegisterResult{}, apiErr(http.StatusBadRequest, code, msg)
}
if strings.EqualFold(req.ID, "inline") {
return RegisterResult{}, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "id reserved")
}
id := req.ID
loginPassword := req.LoginPassword
+76
View File
@@ -141,6 +141,33 @@ func openTestEnv(t *testing.T) *testEnv {
return env
}
func (e *testEnv) setEnabledOnly(t *testing.T, enabled bool) {
t.Helper()
en := "0"
if enabled {
en = "1"
}
now := time.Now().UnixMilli()
err := e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
if _, err := tx.Exec(`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at`,
settingRegistrationEnabled, en, now); err != nil {
return err
}
_, err := tx.Exec(`DELETE FROM settings WHERE key = ?`, settingRegistrationCode)
return err
})
if err != nil {
t.Fatal(err)
}
}
func (l *registerIPLocker) failCount(ip string) int {
l.mu.Lock()
defer l.mu.Unlock()
return len(l.fails[ip])
}
func (e *testEnv) setRegistration(t *testing.T, enabled bool, code string) {
t.Helper()
en := "0"
@@ -233,6 +260,43 @@ func TestRegisterF23_ClosedFails(t *testing.T) {
}
}
func TestRegister_EnabledWithoutUsableCodeIsClosed(t *testing.T) {
cases := []struct {
name string
prep func(*testEnv, *testing.T)
}{
{"no_code_row", func(env *testEnv, t *testing.T) { env.setEnabledOnly(t, true) }},
{"empty_code", func(env *testEnv, t *testing.T) { env.setRegistration(t, true, "") }},
{"short_code", func(env *testEnv, t *testing.T) { env.setRegistration(t, true, "short") }},
}
bodies := []string{
`{"registration_code":"","id":"ep_empty","login_password":"password1"}`,
`{"registration_code":"anything1","id":"ep_any","login_password":"password1"}`,
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
env := openTestEnv(t)
tc.prep(env, t)
for _, body := range bodies {
code, resp, _ := env.doRegister(t, body)
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationClosed {
t.Fatalf("body=%s status=%d resp=%+v", body, code, resp)
}
}
if env.locks.failCount(env.fixedIP) != 0 {
t.Fatalf("unusable code must not count as lock fail: %d", env.locks.failCount(env.fixedIP))
}
logged := env.logBuf.String()
if strings.Contains(logged, "anything1") || strings.Contains(logged, "short") {
t.Fatalf("log leaked code: %s", logged)
}
if !strings.Contains(logged, "unusable_code") {
t.Fatalf("missing warn: %s", logged)
}
})
}
}
func TestRegisterF23_WrongCodeFails_RightCodeOK(t *testing.T) {
env := openTestEnv(t)
env.setRegistration(t, true, "good-code-01")
@@ -475,6 +539,18 @@ func TestRegister_LogOmitsSecrets(t *testing.T) {
}
}
func TestRegister_ReservedInlineID(t *testing.T) {
env := openTestEnv(t)
env.setRegistration(t, true, "inline-code1")
code, resp, _ := env.doRegister(t, `{"registration_code":"inline-code1","id":"inline","login_password":"password1"}`)
if code != http.StatusBadRequest || resp.Error == nil || resp.Error.Code != protocol.CodeBadRequest {
t.Fatalf("status=%d resp=%+v", code, resp)
}
if _, _, ok := env.getEndpoint(t, "inline"); ok {
t.Fatal("inline must not be created")
}
}
func TestRegister_NSTPasswordRejected(t *testing.T) {
env := openTestEnv(t)
env.setRegistration(t, true, "nst-code-01")
+5 -3
View File
@@ -122,13 +122,15 @@ WHERE id > ?
ORDER BY id ASC
LIMIT ?`, cursor, limit+1)
} else {
like := "%" + strings.ToLower(query) + "%"
prefix := strings.ToLower(query) + "%"
q := strings.ToLower(query)
esc := escapeLikePattern(q)
like := "%" + esc + "%"
prefix := esc + "%"
rows, err = a.db.Read.QueryContext(ctx, `
SELECT id, name, online_since, offline_since, talk_hash
FROM endpoints
WHERE id > ?
AND (lower(id) LIKE ? OR lower(name) LIKE ?)
AND (lower(id) LIKE ? ESCAPE '\' OR lower(name) LIKE ? ESCAPE '\')
ORDER BY id ASC
LIMIT ?`, cursor, prefix, like, limit+1)
}
+16
View File
@@ -0,0 +1,16 @@
package presence
import "strings"
func escapeLikePattern(s string) string {
var b strings.Builder
b.Grow(len(s) + 4)
for _, r := range s {
switch r {
case '\\', '%', '_':
b.WriteByte('\\')
}
b.WriteRune(r)
}
return b.String()
}
+29
View File
@@ -131,6 +131,35 @@ func TestF03PresenceAndDirectory(t *testing.T) {
}
}
func TestDirectoryLikeEscapesUnderscore(t *testing.T) {
t.Parallel()
app, db, _ := openPresence(t)
ctx := context.Background()
insertEP(t, db, "e_ab1", "underscore")
insertEP(t, db, "exab2", "wildcard")
insertEP(t, db, "pct", "has%percent")
q, _, err := app.Directory(ctx, &protocol.DirectoryList{
V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "u1", Query: "e_ab", Limit: 10,
})
if err != nil {
t.Fatal(err)
}
if len(q) != 1 || q[0].ID != "e_ab1" {
t.Fatalf("want only e_ab1, got %+v", q)
}
q, _, err = app.Directory(ctx, &protocol.DirectoryList{
V: protocol.Version, Type: protocol.TypeDirectoryList, RID: "u2", Query: "%", Limit: 10,
})
if err != nil {
t.Fatal(err)
}
if len(q) != 1 || q[0].ID != "pct" {
t.Fatalf("literal %% should not match all, got %+v", q)
}
}
func TestF04PresenceWatch(t *testing.T) {
t.Parallel()
app, db, down := openPresence(t)
+121 -8
View File
@@ -21,23 +21,34 @@ func policyFor(kind LockKind) lockPolicy {
}
}
const (
defaultMaxLockEntries = 65536
lockSweepEveryOps = 1024
lockSweepInterval = time.Minute
)
type lockEntry struct {
fails []time.Time
lockedUntil time.Time
lastFail time.Time
}
// MemoryLocks 是内存锁定计数器(重启清零)。
type MemoryLocks struct {
mu sync.Mutex
entries map[string]*lockEntry
now func() time.Time
mu sync.Mutex
entries map[string]*lockEntry
now func() time.Time
ops int
lastSweep time.Time
maxEntries int
}
// NewLoginLocks 创建默认锁定计数器。
func NewLoginLocks() *MemoryLocks {
return &MemoryLocks{
entries: make(map[string]*lockEntry),
now: time.Now,
entries: make(map[string]*lockEntry),
now: time.Now,
maxEntries: defaultMaxLockEntries,
}
}
@@ -56,23 +67,121 @@ func lockMapKey(key LockKey) string {
return string(key.Kind) + "|" + key.EndpointID + "|" + key.PeerID + "|" + key.IP
}
func (l *MemoryLocks) maxCap() int {
if l.maxEntries <= 0 {
return defaultMaxLockEntries
}
return l.maxEntries
}
func (e *lockEntry) stale(now time.Time, window time.Duration) bool {
if e == nil {
return true
}
if e.lockedUntil.After(now) {
return false
}
if e.lastFail.IsZero() {
return true
}
return !e.lastFail.After(now.Add(-window))
}
func (l *MemoryLocks) maybeSweepLocked(now time.Time) {
max := l.maxCap()
if l.lastSweep.IsZero() {
l.lastSweep = now
}
l.ops++
if l.ops < lockSweepEveryOps && now.Sub(l.lastSweep) < lockSweepInterval && len(l.entries) <= max {
return
}
l.ops = 0
l.lastSweep = now
l.sweepExpiredLocked(now)
l.enforceCapLocked(now)
}
func (l *MemoryLocks) sweepExpiredLocked(now time.Time) {
for k, e := range l.entries {
window := policyFor("").window
parts := splitLockKey(k)
if len(parts) == 4 {
window = policyFor(LockKind(parts[0])).window
}
if e.stale(now, window) {
delete(l.entries, k)
}
}
}
func (l *MemoryLocks) enforceCapLocked(now time.Time) {
max := l.maxCap()
for len(l.entries) > max {
var (
victim string
victimLocked bool
victimTime time.Time
found bool
)
for k, e := range l.entries {
locked := e.lockedUntil.After(now)
t := e.lastFail
if locked {
t = e.lockedUntil
}
better := !found
if found {
if victimLocked && !locked {
better = true
} else if victimLocked == locked && (t.Before(victimTime) || (t.Equal(victimTime) && k < victim)) {
better = true
}
}
if better {
found = true
victim, victimLocked, victimTime = k, locked, t
}
}
if !found {
return
}
delete(l.entries, victim)
}
}
func (l *MemoryLocks) entryCount() int {
l.mu.Lock()
defer l.mu.Unlock()
return len(l.entries)
}
// Check 若当前已锁定返回 locked=true 与剩余时间。
func (l *MemoryLocks) Check(key LockKey) (bool, time.Duration) {
l.mu.Lock()
defer l.mu.Unlock()
now := l.now()
e := l.entries[lockMapKey(key)]
k := lockMapKey(key)
e := l.entries[k]
if e == nil {
l.maybeSweepLocked(now)
return false, 0
}
if e.lockedUntil.After(now) {
l.maybeSweepLocked(now)
return true, e.lockedUntil.Sub(now)
}
// 到期自动解除:清空锁定与窗口内失败(保留结构以便后续 Fail)。
if !e.lockedUntil.IsZero() && !e.lockedUntil.After(now) {
if e.stale(now, policyFor(key.Kind).window) {
delete(l.entries, k)
l.maybeSweepLocked(now)
return false, 0
}
// 到期自动解除:清空锁定与窗口内失败(保留仍在窗口内的失败计数)。
if !e.lockedUntil.IsZero() {
e.lockedUntil = time.Time{}
e.fails = nil
}
l.maybeSweepLocked(now)
return false, 0
}
@@ -88,6 +197,7 @@ func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) {
l.entries[k] = e
}
if e.lockedUntil.After(now) {
l.maybeSweepLocked(now)
return true, e.lockedUntil.Sub(now)
}
if !e.lockedUntil.IsZero() {
@@ -103,11 +213,14 @@ func (l *MemoryLocks) Fail(key LockKey) (bool, time.Duration) {
}
}
e.fails = append(kept, now)
e.lastFail = now
if len(e.fails) >= pol.maxFails {
e.lockedUntil = now.Add(pol.lockFor)
e.fails = nil
l.maybeSweepLocked(now)
return true, pol.lockFor
}
l.maybeSweepLocked(now)
return false, 0
}
+60
View File
@@ -0,0 +1,60 @@
package auth
import (
"fmt"
"testing"
"time"
)
func TestLoginLocksSweepExpiredKeepsLocked(t *testing.T) {
locks := NewLoginLocks()
now := time.Date(2026, 9, 30, 12, 0, 0, 0, time.UTC)
locks.SetClock(func() time.Time { return now })
for i := 0; i < 10000; i++ {
if locked, _ := locks.Fail(LockKey{Kind: LockRegisterIP, IP: fmt.Sprintf("2001:db8::%d", i)}); locked {
t.Fatalf("unexpected lock at i=%d", i)
}
}
if n := locks.entryCount(); n < 10000 {
t.Fatalf("want 10000 entries before sweep, got %d", n)
}
now = now.Add(4 * time.Minute)
keep := LockKey{Kind: LockRegisterIP, IP: "keep-locked"}
for i := 0; i < 10; i++ {
locks.Fail(keep)
}
if locked, _ := locks.Check(keep); !locked {
t.Fatal("keep-locked should be locked")
}
now = now.Add(2 * time.Minute) // 10k 已过 5 分钟窗口;keep 仍锁定至 +5min
if locked, _ := locks.Check(LockKey{Kind: LockRegisterIP, IP: "trigger-sweep"}); locked {
t.Fatal("trigger must not lock")
}
if n := locks.entryCount(); n != 1 {
t.Fatalf("want only locked entry after sweep, got %d", n)
}
if locked, _ := locks.Check(keep); !locked {
t.Fatal("locked entry must remain")
}
}
func TestLoginLocksCapDoesNotGrow(t *testing.T) {
locks := NewLoginLocks()
locks.maxEntries = 64
now := time.Date(2026, 9, 30, 15, 0, 0, 0, time.UTC)
locks.SetClock(func() time.Time { return now })
for i := 0; i < 200; i++ {
locks.Fail(LockKey{Kind: LockRegisterIP, IP: fmt.Sprintf("ip-%d", i)})
now = now.Add(time.Millisecond)
if n := locks.entryCount(); n > 64 {
t.Fatalf("entries=%d exceeded cap at i=%d", n, i)
}
}
if n := locks.entryCount(); n != 64 {
t.Fatalf("entries=%d want 64", n)
}
}