Files
NixMsg/internal/app/identity/register_test.go
T

602 lines
20 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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 &registerIPLocker{
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) ClearAllForEndpoint(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_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")
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)
}
}