447 lines
14 KiB
Go
447 lines
14 KiB
Go
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/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) 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 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")
|
||
}
|
||
}
|
||
|
||
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)
|
||
}
|
||
}
|