package identity import ( "bytes" "context" "database/sql" "encoding/json" "io" "log/slog" "net/http" "net/http/httptest" "strings" "sync" "testing" "time" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/httpx" "git.asio.asia/nixevol/NixMsg/internal/protocol" "git.asio.asia/nixevol/NixMsg/internal/store" ) // registerIPLocker 仅实现 LockRegisterIP:5 分钟窗口内 10 次失败则锁定 5 分钟。 // P3 的完整 LoginLocks 未合入本分支时,测试用此可替换实现覆盖 F23 锁定验收。 type registerIPLocker struct { mu sync.Mutex fails map[string][]time.Time lockedUntil map[string]time.Time now func() time.Time window time.Duration limit int lockFor time.Duration } func newRegisterIPLocker(now func() time.Time) *registerIPLocker { if now == nil { now = time.Now } return ®isterIPLocker{ fails: make(map[string][]time.Time), lockedUntil: make(map[string]time.Time), now: now, window: 5 * time.Minute, limit: 10, lockFor: 5 * time.Minute, } } func (l *registerIPLocker) Check(key auth.LockKey) (bool, time.Duration) { if key.Kind != auth.LockRegisterIP { return false, 0 } l.mu.Lock() defer l.mu.Unlock() until, ok := l.lockedUntil[key.IP] if !ok { return false, 0 } now := l.now() if now.Before(until) { return true, until.Sub(now) } delete(l.lockedUntil, key.IP) return false, 0 } func (l *registerIPLocker) Fail(key auth.LockKey) (bool, time.Duration) { if key.Kind != auth.LockRegisterIP { return false, 0 } l.mu.Lock() defer l.mu.Unlock() now := l.now() if until, ok := l.lockedUntil[key.IP]; ok && now.Before(until) { return true, until.Sub(now) } cutoff := now.Add(-l.window) list := l.fails[key.IP] kept := list[:0] for _, t := range list { if t.After(cutoff) { kept = append(kept, t) } } kept = append(kept, now) l.fails[key.IP] = kept if len(kept) >= l.limit { until := now.Add(l.lockFor) l.lockedUntil[key.IP] = until return true, l.lockFor } return false, 0 } func (l *registerIPLocker) ClearEndpoint(string) {} func (l *registerIPLocker) Clear(key auth.LockKey) { l.mu.Lock() defer l.mu.Unlock() delete(l.fails, key.IP) delete(l.lockedUntil, key.IP) } var _ auth.LoginLocks = (*registerIPLocker)(nil) type testEnv struct { db *store.DB hash auth.HashPool locks *registerIPLocker logBuf *bytes.Buffer handler http.Handler fixedIP string } func openTestEnv(t *testing.T) *testEnv { t.Helper() db, err := store.Open(t.TempDir(), "FULL") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = db.Close() }) buf := &bytes.Buffer{} logger := slog.New(slog.NewTextHandler(buf, &slog.HandlerOptions{Level: slog.LevelInfo})) locks := newRegisterIPLocker(time.Now) env := &testEnv{ db: db, hash: auth.NewStubHashPool(), locks: locks, logBuf: buf, fixedIP: "203.0.113.10", } env.handler = NewRegisterHandler(RegisterConfig{ DB: db, Hash: env.hash, Locks: locks, Logger: logger, ClientIP: func(*http.Request) string { return env.fixedIP }, }) 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" 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(`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?) ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=excluded.updated_at`, settingRegistrationCode, code, now) return err }) if err != nil { t.Fatal(err) } } func (e *testEnv) insertEndpoint(t *testing.T, id, loginHash string) { t.Helper() err := e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error { _, err := tx.Exec(`INSERT INTO endpoints( id, name, remark, source, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at ) VALUES (?, '', '', 'admin', ?, NULL, 0, 0, 1, ?)`, id, loginHash, time.Now().UnixMilli()) return err }) if err != nil { t.Fatal(err) } } func (e *testEnv) getEndpoint(t *testing.T, id string) (source, loginHash string, ok bool) { t.Helper() err := e.db.Read.QueryRow(`SELECT source, login_hash FROM endpoints WHERE id = ?`, id).Scan(&source, &loginHash) if errorsIsNoRows(err) { return "", "", false } if err != nil { t.Fatal(err) } return source, loginHash, true } func errorsIsNoRows(err error) bool { return err == sql.ErrNoRows } type registerResp struct { OK bool `json:"ok"` Data struct { ID string `json:"id"` LoginPassword string `json:"login_password"` } `json:"data"` Error *protocol.ErrorBody `json:"error"` } func (e *testEnv) doRegister(t *testing.T, body string) (int, registerResp, http.Header) { t.Helper() req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body)) req.Header.Set("Content-Type", "application/json") req.RemoteAddr = e.fixedIP + ":54321" rr := httptest.NewRecorder() e.handler.ServeHTTP(rr, req) var resp registerResp if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil { t.Fatalf("decode resp: %v body=%s", err, rr.Body.String()) } return rr.Code, resp, rr.Header() } func TestRegisterF23_ClosedFails(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, false, "secretcode") code, resp, hdr := env.doRegister(t, `{"registration_code":"secretcode","id":"ep_closed","login_password":"password1"}`) if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationClosed { t.Fatalf("status=%d resp=%+v", code, resp) } if hdr.Get("Access-Control-Allow-Origin") != "*" { t.Fatalf("missing CORS: %v", hdr) } if _, _, ok := env.getEndpoint(t, "ep_closed"); ok { t.Fatal("endpoint should not be created when closed") } } 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") code, resp, _ := env.doRegister(t, `{"registration_code":"bad-code-xx","id":"ep_wrong","login_password":"password1"}`) if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid { t.Fatalf("wrong code: status=%d resp=%+v", code, resp) } code, resp, hdr := env.doRegister(t, `{"registration_code":"good-code-01","id":"ep_ok1","login_password":"password1","name":"门口"}`) if code != http.StatusOK || !resp.OK || resp.Data.ID != "ep_ok1" { t.Fatalf("ok register: status=%d resp=%+v", code, resp) } if resp.Data.LoginPassword != "" { t.Fatalf("provided password must not echo: %q", resp.Data.LoginPassword) } if hdr.Get("Access-Control-Allow-Origin") != "*" { t.Fatal("missing CORS on success") } source, loginHash, ok := env.getEndpoint(t, "ep_ok1") if !ok || source != "self" { t.Fatalf("endpoint source=%q ok=%v", source, ok) } match, err := env.hash.Verify(context.Background(), auth.PasswordLogin, "password1", loginHash) if err != nil || !match { t.Fatalf("login hash verify: match=%v err=%v", match, err) } } func TestRegisterF23_ChangeCode_OldFails_ExistingRemains(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "code-old-01") code, resp, _ := env.doRegister(t, `{"registration_code":"code-old-01","id":"ep_keep","login_password":"password1"}`) if code != http.StatusOK || resp.Data.ID != "ep_keep" { t.Fatalf("first register: status=%d resp=%+v", code, resp) } _, oldHash, ok := env.getEndpoint(t, "ep_keep") if !ok { t.Fatal("missing endpoint after register") } env.setRegistration(t, true, "code-new-02") code, resp, _ = env.doRegister(t, `{"registration_code":"code-old-01","id":"ep_new","login_password":"password1"}`) if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid { t.Fatalf("old code after rotate: status=%d resp=%+v", code, resp) } code, resp, _ = env.doRegister(t, `{"registration_code":"code-new-02","id":"ep_new","login_password":"password1"}`) if code != http.StatusOK || resp.Data.ID != "ep_new" { t.Fatalf("new code: status=%d resp=%+v", code, resp) } _, hashAfter, ok := env.getEndpoint(t, "ep_keep") if !ok || hashAfter != oldHash { t.Fatalf("existing endpoint mutated: ok=%v hashEqual=%v", ok, hashAfter == oldHash) } } func TestRegisterF23_WrongCodeLock(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "lock-code-1") for i := 0; i < 10; i++ { code, resp, _ := env.doRegister(t, `{"registration_code":"wrong-code","id":"ep_lock","login_password":"password1"}`) if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid { t.Fatalf("fail #%d: status=%d resp=%+v", i+1, code, resp) } } code, resp, _ := env.doRegister(t, `{"registration_code":"lock-code-1","id":"ep_lock","login_password":"password1"}`) if code != http.StatusTooManyRequests || resp.Error == nil || resp.Error.Code != protocol.CodeRateLimited { t.Fatalf("locked with good code: status=%d resp=%+v", code, resp) } if _, _, ok := env.getEndpoint(t, "ep_lock"); ok { t.Fatal("must not insert while rate limited") } } // TestRegisterTrustedProxyClientIPLock 验证与管理接口相同的 httpx.ClientIP: // 受信代理的 X-Forwarded-For 按真实客户端 IP 计锁;非信任来源不采信转发头。 func TestRegisterTrustedProxyClientIPLock(t *testing.T) { db, err := store.Open(t.TempDir(), "FULL") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = db.Close() }) locks := newRegisterIPLocker(time.Now) trusted := httpx.ParseCIDRs([]string{"127.0.0.1/32"}) handler := NewRegisterHandler(RegisterConfig{ DB: db, Hash: auth.NewStubHashPool(), Locks: locks, ClientIP: func(r *http.Request) string { return httpx.ClientIP(r, trusted) }, }) env := &testEnv{db: db, hash: auth.NewStubHashPool(), locks: locks, handler: handler} env.setRegistration(t, true, "proxy-lock-1") post := func(remote, xff, body string) (int, registerResp) { t.Helper() req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body)) req.Header.Set("Content-Type", "application/json") req.RemoteAddr = remote if xff != "" { req.Header.Set("X-Forwarded-For", xff) } rr := httptest.NewRecorder() handler.ServeHTTP(rr, req) var resp registerResp if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil { t.Fatalf("decode: %v body=%s", err, rr.Body.String()) } return rr.Code, resp } wrong := `{"registration_code":"wrong-code","id":"ep_px","login_password":"password1"}` good := `{"registration_code":"proxy-lock-1","id":"ep_px","login_password":"password1"}` for i := 0; i < 10; i++ { code, resp := post("127.0.0.1:9000", "198.51.100.7", wrong) if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid { t.Fatalf("trusted fail #%d: status=%d resp=%+v", i+1, code, resp) } } code, resp := post("127.0.0.1:9000", "198.51.100.7", good) if code != http.StatusTooManyRequests || resp.Error == nil || resp.Error.Code != protocol.CodeRateLimited { t.Fatalf("real client should be locked: status=%d resp=%+v", code, resp) } code, resp = post("127.0.0.1:9000", "198.51.100.8", good) if code != http.StatusOK || !resp.OK || resp.Data.ID != "ep_px" { t.Fatalf("other XFF client must not share lock: status=%d resp=%+v", code, resp) } // 非信任对端:忽略 XFF,按 RemoteAddr 计锁。 locks.Clear(auth.LockKey{Kind: auth.LockRegisterIP, IP: "198.51.100.7"}) locks.Clear(auth.LockKey{Kind: auth.LockRegisterIP, IP: "203.0.113.50"}) for i := 0; i < 10; i++ { code, resp = post("203.0.113.50:4433", "198.51.100.7", wrong) if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid { t.Fatalf("untrusted fail #%d: status=%d resp=%+v", i+1, code, resp) } } code, resp = post("203.0.113.50:4433", "198.51.100.7", `{"registration_code":"proxy-lock-1","id":"ep_px2","login_password":"password1"}`) if code != http.StatusTooManyRequests || resp.Error == nil || resp.Error.Code != protocol.CodeRateLimited { t.Fatalf("untrusted RemoteAddr should be locked: status=%d resp=%+v", code, resp) } // 若误采信 XFF,198.51.100.7 会已锁;直连该 IP 应仍可注册。 code, resp = post("198.51.100.7:5555", "", `{"registration_code":"proxy-lock-1","id":"ep_px3","login_password":"password1"}`) if code != http.StatusOK || !resp.OK || resp.Data.ID != "ep_px3" { t.Fatalf("spoofed XFF must not lock real client: status=%d resp=%+v", code, resp) } } func TestRegisterF23_IDTakenKeepsOriginal(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "taken-code") env.insertEndpoint(t, "ep_taken", "stub$original-password-xx") code, resp, _ := env.doRegister(t, `{"registration_code":"taken-code","id":"ep_taken","login_password":"password1"}`) if code != http.StatusConflict || resp.Error == nil || resp.Error.Code != protocol.CodeIDTaken { t.Fatalf("id taken: status=%d resp=%+v", code, resp) } source, loginHash, ok := env.getEndpoint(t, "ep_taken") if !ok || source != "admin" || loginHash != "stub$original-password-xx" { t.Fatalf("original endpoint changed: source=%q hash=%q", source, loginHash) } } func TestRegister_GenerateIDAndPassword(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "gen-code-01") code, resp, _ := env.doRegister(t, `{"registration_code":"gen-code-01","id":"","login_password":""}`) if code != http.StatusOK || !resp.OK { t.Fatalf("status=%d resp=%+v", code, resp) } if !strings.HasPrefix(resp.Data.ID, "e_") || len(resp.Data.ID) != 10 { t.Fatalf("generated id=%q", resp.Data.ID) } if len(resp.Data.LoginPassword) < protocol.MinLoginPasswordLen { t.Fatalf("generated password too short: %q", resp.Data.LoginPassword) } if strings.HasPrefix(resp.Data.LoginPassword, protocol.SessionTokenPrefix) { t.Fatal("generated password starts with nst_") } source, _, ok := env.getEndpoint(t, resp.Data.ID) if !ok || source != "self" { t.Fatalf("source=%q ok=%v", source, ok) } } func TestRegister_OPTIONS_CORS(t *testing.T) { env := openTestEnv(t) req := httptest.NewRequest(http.MethodOptions, "/api/client/register", nil) rr := httptest.NewRecorder() env.handler.ServeHTTP(rr, req) if rr.Code != http.StatusNoContent { t.Fatalf("status=%d", rr.Code) } if rr.Header().Get("Access-Control-Allow-Origin") != "*" { t.Fatal("missing Allow-Origin") } if !strings.Contains(rr.Header().Get("Access-Control-Allow-Methods"), "POST") { t.Fatalf("methods=%q", rr.Header().Get("Access-Control-Allow-Methods")) } if len(rr.Result().Cookies()) != 0 { t.Fatal("must not set cookies") } } func TestRegister_BodyTooLarge(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "big-code-01") body := `{"registration_code":"big-code-01","id":"ep_big","login_password":"password1","name":"` + strings.Repeat("x", 5000) + `"}` req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body)) req.Header.Set("Content-Type", "application/json") rr := httptest.NewRecorder() env.handler.ServeHTTP(rr, req) if rr.Code != http.StatusBadRequest { t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String()) } var resp registerResp _ = json.Unmarshal(rr.Body.Bytes(), &resp) if resp.Error == nil || resp.Error.Code != protocol.CodeBadRequest { t.Fatalf("resp=%+v", resp) } } func TestRegister_LogOmitsSecrets(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "log-secret-code") _, _, _ = env.doRegister(t, `{"registration_code":"log-secret-code","id":"ep_log","login_password":"supersecretpw"}`) logged := env.logBuf.String() if strings.Contains(logged, "log-secret-code") || strings.Contains(logged, "supersecretpw") { t.Fatalf("log leaked secrets: %s", logged) } if !strings.Contains(logged, "ep_log") || !strings.Contains(logged, env.fixedIP) { t.Fatalf("log missing id/ip: %s", logged) } } func TestRegister_NSTPasswordRejected(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "nst-code-01") code, resp, _ := env.doRegister(t, `{"registration_code":"nst-code-01","id":"ep_nst","login_password":"nst_notallowed"}`) if code != http.StatusBadRequest || resp.Error == nil || resp.Error.Code != protocol.CodeBadRequest { t.Fatalf("status=%d resp=%+v", code, resp) } } func TestMountOnServeMux(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "mux-code-01") mux := http.NewServeMux() mux.Handle("/api/client/register", env.handler) body := `{"registration_code":"mux-code-01","id":"ep_mux","login_password":"password1"}` req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body)) rr := httptest.NewRecorder() mux.ServeHTTP(rr, req) if rr.Code != http.StatusOK { t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String()) } } func TestServerRegisterInterface(t *testing.T) { env := openTestEnv(t) env.setRegistration(t, true, "svc-code-01") svc := NewServer(RegisterConfig{ DB: env.db, Hash: env.hash, Locks: env.locks, Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), }) res, err := svc.Register(context.Background(), RegisterRequest{ RegistrationCode: "svc-code-01", ID: "ep_svc", LoginPassword: "password1", RemoteIP: "198.51.100.1", Source: "self", }) if err != nil { t.Fatal(err) } if res.ID != "ep_svc" { t.Fatalf("id=%q", res.ID) } }