From 308b0b9edd3e56d2f669d65a14817670e6ccf419 Mon Sep 17 00:00:00 2001 From: Nixevol Date: Wed, 30 Sep 2026 06:47:47 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E5=AE=9E=E7=8E=B0=E7=AB=AF=E8=87=AA?= =?UTF-8?q?=E5=8A=A9=E6=B3=A8=E5=86=8C=20HTTP=20=E6=8E=A5=E5=8F=A3?= =?UTF-8?q?=E4=B8=8E=20F23=20=E9=AA=8C=E6=94=B6=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/DEVIATIONS.md | 37 +- internal/app/identity/register.go | 417 +++++++++++++++++++++++ internal/app/identity/register_test.go | 446 +++++++++++++++++++++++++ 3 files changed, 899 insertions(+), 1 deletion(-) create mode 100644 internal/app/identity/register.go create mode 100644 internal/app/identity/register_test.go diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index 1d2e76e..337f802 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -310,7 +310,42 @@ ## 身份 I -暂无。 +### I1 2026-09-30 + +1. **注册做成可挂载 Handler,不改 cmd/listener** + - 原条款:TASKS I1 / DEVELOPMENT 6.9 在 `listen` 上提供 `POST /api/client/register`;依赖 N1 端口识别。 + - 实际做法:`identity.NewRegisterHandler` / `identity.NewServer().Handler()` 返回 `http.Handler`,由接线方 `mux.Handle("/api/client/register", h)`;本线不改 `cmd/nixmsg`、`internal/listener`(N 线未合入)。 + - 原因:隔离交付,避免抢 N/P 接线。 + - 备选方案:本线直接改 `wire.go` 挂路由。 + - 影响:合入后需总控或 N/A 接线才对外可访问。 + +2. **密码哈希与锁定走 auth 接口,本分支用可替换假实现测** + - 原条款:依赖 P3 argon2 池与锁定计数器。 + - 实际做法:`RegisterConfig.Hash`/`Locks` 注入 `auth.HashPool`、`auth.LoginLocks`;测试用 `auth.NewStubHashPool` + 仅实现 `LockRegisterIP`(5 分钟 10 次)的测试锁定器,不在本线重写 argon2。 + - 原因:P3 尚未在本分支。 + - 备选方案:等 P3 合入后再写 I1。 + - 影响:生产须注入 P3 实现;StubLoginLocks 永不锁定,不能直接用于开放注册。 + +3. **settings 开关取值** + - 原条款:`settings.registration_enabled`,未规定字符串字面量。 + - 实际做法:`1`/`true`/`yes`/`on`(大小写不敏感)视为开启,其余(含缺省)关闭;安全码键 `registration_code`。 + - 原因:与 store 测试写入的 `"0"`/`"1"` 对齐并兼容常见布尔字面量。 + - 备选方案:仅认 `"1"`。 + - 影响:A 线写注册设置时宜写 `"1"`/`"0"`。 + +4. **客户端 IP** + - 原条款:DEVELOPMENT 4.5 受信任代理下用 `X-Forwarded-For`。 + - 实际做法:Handler 默认取 `RemoteAddr` 的 host;可通过 `RegisterConfig.ClientIP` 注入。本线不做 `trusted_proxies` 解析(属 listener/接线)。 + - 原因:不改 listener;代理 IP 应由外层在挂载前算好或注入。 + - 备选方案:在 identity 内复制 4.5 逻辑。 + - 影响:经代理部署时接线方必须注入真实 IP,否则锁定按直连 IP 计。 + +5. **生成登录密码长度** + - 原条款:F01 留空则生成,8–128 字符,不以 `nst_` 开头;未规定生成长度。 + - 实际做法:生成 20 位字母数字;若偶然以 `nst_` 开头则重抽。 + - 原因:与管理员 init 量级接近,满足规则。 + - 备选方案:16/32 位。 + - 影响:无产品行为差异。 ## 后台接口 A diff --git a/internal/app/identity/register.go b/internal/app/identity/register.go new file mode 100644 index 0000000..020a500 --- /dev/null +++ b/internal/app/identity/register.go @@ -0,0 +1,417 @@ +package identity + +import ( + "context" + "crypto/rand" + "crypto/subtle" + "database/sql" + "errors" + "io" + "log/slog" + "net" + "net/http" + "strings" + "time" + "unicode/utf8" + + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/protocol" + "git.asio.asia/nixevol/NixMsg/internal/store" +) + +const ( + maxRegisterBodyBytes = 4 * 1024 + + settingRegistrationEnabled = "registration_enabled" + settingRegistrationCode = "registration_code" + + sourceSelf = "self" + + idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789" +) + +// APIError 是注册 HTTP/业务错误,带 HTTP 状态与协议错误码。 +type APIError struct { + Status int + Code string + Message string +} + +func (e *APIError) Error() string { + if e == nil { + return "" + } + if e.Message == "" { + return e.Code + } + return e.Code + ": " + e.Message +} + +func apiErr(status int, code, msg string) *APIError { + return &APIError{Status: status, Code: code, Message: msg} +} + +// RegisterConfig 是可挂载注册处理器的依赖。 +// Hash / Locks 用 auth 接口;P3 未合入时测试可注入 StubHashPool 与可锁定的 LoginLocks。 +type RegisterConfig struct { + DB *store.DB + Hash auth.HashPool + Locks auth.LoginLocks + Logger *slog.Logger + // Now 可测;nil 则用 time.Now。 + Now func() time.Time + // ClientIP 可测;nil 则从 RemoteAddr 取 host。 + ClientIP func(*http.Request) string +} + +// RegisterHandler 处理 POST/OPTIONS /api/client/register(可挂到任意 ServeMux)。 +type RegisterHandler struct { + cfg RegisterConfig +} + +// NewRegisterHandler 构造可挂载的注册 Handler。DB/Hash/Locks 必填。 +func NewRegisterHandler(cfg RegisterConfig) *RegisterHandler { + if cfg.Logger == nil { + cfg.Logger = slog.Default() + } + if cfg.Now == nil { + cfg.Now = time.Now + } + if cfg.ClientIP == nil { + cfg.ClientIP = clientIPFromRemoteAddr + } + return &RegisterHandler{cfg: cfg} +} + +// ServeHTTP 实现 http.Handler。 +func (h *RegisterHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + setCORS(w) + switch r.Method { + case http.MethodOptions: + w.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS") + w.Header().Set("Access-Control-Allow-Headers", "Content-Type") + w.WriteHeader(http.StatusNoContent) + return + case http.MethodPost: + h.handlePost(w, r) + default: + writeRegisterError(w, apiErr(http.StatusMethodNotAllowed, protocol.CodeBadRequest, "method not allowed")) + } +} + +func (h *RegisterHandler) handlePost(w http.ResponseWriter, r *http.Request) { + ip := h.cfg.ClientIP(r) + r.Body = http.MaxBytesReader(w, r.Body, maxRegisterBodyBytes) + body, err := io.ReadAll(r.Body) + if err != nil { + var maxErr *http.MaxBytesError + if errors.As(err, &maxErr) || errors.Is(err, io.ErrUnexpectedEOF) || isBodyTooLarge(err) { + h.logResult("bad_request", "", ip) + writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "body too large")) + return + } + h.logResult("bad_request", "", ip) + writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "read body failed")) + return + } + + req, err := protocol.DecodeRegister(body) + if err != nil { + h.logResult("bad_request", "", ip) + writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "invalid json")) + return + } + + result, apiErr := h.register(r.Context(), req, ip) + if apiErr != nil { + id := "" + if req != nil { + id = req.ID + } + h.logResult(apiErr.Code, id, ip) + writeRegisterError(w, apiErr) + return + } + h.logResult("ok", result.ID, ip) + writeRegisterOK(w, result) +} + +func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRequest, ip string) (RegisterResult, *APIError) { + enabled, storedCode, err := h.loadRegistrationSettings(ctx) + if err != nil { + return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "settings unavailable") + } + if !enabled { + 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 { + return RegisterResult{}, apiErr(http.StatusTooManyRequests, protocol.CodeRateLimited, "rate limited") + } + + if !constantTimeEqual(req.RegistrationCode, storedCode) { + h.cfg.Locks.Fail(lockKey) + return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationCodeInvalid, "registration code invalid") + } + + if valErr := req.Validate(); valErr != nil { + code, msg := protocol.CodeBadRequest, valErr.Error() + var pe *protocol.Error + if errors.As(valErr, &pe) && pe != nil { + code, msg = pe.Code, pe.Message + } + return RegisterResult{}, apiErr(http.StatusBadRequest, code, msg) + } + + id := req.ID + loginPassword := req.LoginPassword + passwordGenerated := false + if loginPassword == "" { + pw, genErr := generateLoginPassword() + if genErr != nil { + return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate password failed") + } + loginPassword = pw + passwordGenerated = true + } + + loginHash, err := h.cfg.Hash.Hash(ctx, auth.PasswordLogin, loginPassword) + if err != nil { + return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "hash failed") + } + + var talkHash sql.NullString + if req.TalkPassword != "" { + th, hashErr := h.cfg.Hash.Hash(ctx, auth.PasswordTalk, req.TalkPassword) + if hashErr != nil { + return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "hash failed") + } + talkHash = sql.NullString{String: th, Valid: true} + } + + nowMs := h.cfg.Now().UnixMilli() + const maxIDAttempts = 8 + for attempt := 0; attempt < maxIDAttempts; attempt++ { + useID := id + if useID == "" { + genID, genErr := generateEndpointID() + if genErr != nil { + return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate id failed") + } + useID = genID + } + + insertErr := h.cfg.DB.Queue.Do(ctx, func(tx *sql.Tx) error { + _, execErr := tx.ExecContext(ctx, ` +INSERT INTO endpoints( + id, name, remark, source, login_hash, talk_hash, talk_version, + default_delay_ms, enabled, created_at +) VALUES (?, ?, '', ?, ?, ?, 0, 0, 1, ?)`, + useID, req.Name, sourceSelf, loginHash, talkHash, nowMs, + ) + return execErr + }) + if insertErr == nil { + out := RegisterResult{ID: useID} + if passwordGenerated { + out.LoginPassword = loginPassword + } + return out, nil + } + if isUniqueConstraint(insertErr) { + if id != "" { + return RegisterResult{}, apiErr(http.StatusConflict, protocol.CodeIDTaken, "id taken") + } + continue + } + return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "insert failed") + } + return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate id exhausted") +} + +func (h *RegisterHandler) loadRegistrationSettings(ctx context.Context) (enabled bool, code string, err error) { + var enabledVal, codeVal sql.NullString + row := h.cfg.DB.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingRegistrationEnabled) + if scanErr := row.Scan(&enabledVal); scanErr != nil && !errors.Is(scanErr, sql.ErrNoRows) { + return false, "", scanErr + } + row = h.cfg.DB.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingRegistrationCode) + if scanErr := row.Scan(&codeVal); scanErr != nil && !errors.Is(scanErr, sql.ErrNoRows) { + return false, "", scanErr + } + return settingTruthy(enabledVal.String), codeVal.String, nil +} + +func (h *RegisterHandler) logResult(result, id, ip string) { + h.cfg.Logger.Info("register", "result", result, "id", id, "ip", ip) +} + +// Server 实现 identity.Service:I1 只实现 Register,其余仍为未实现。 +type Server struct { + handler *RegisterHandler +} + +// NewServer 用同一套依赖构造 Service(Register)与可挂载 Handler。 +func NewServer(cfg RegisterConfig) *Server { + return &Server{handler: NewRegisterHandler(cfg)} +} + +// Handler 返回可挂载的注册 HTTP 处理器。 +func (s *Server) Handler() http.Handler { return s.handler } + +// Register 实现自助注册(source 固定为 self;RemoteIP 用于锁定)。 +func (s *Server) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) { + preq := &protocol.RegisterRequest{ + RegistrationCode: req.RegistrationCode, + ID: req.ID, + LoginPassword: req.LoginPassword, + Name: req.Name, + TalkPassword: req.TalkPassword, + } + ip := req.RemoteIP + if ip == "" { + ip = "0.0.0.0" + } + result, err := s.handler.register(ctx, preq, ip) + if err != nil { + return RegisterResult{}, err + } + return result, nil +} + +func (s *Server) SelfGet(context.Context, string) (SelfInfo, error) { + return SelfInfo{}, ErrNotImplemented +} +func (s *Server) SelfUpdate(context.Context, string, *protocol.SelfUpdate) error { + return ErrNotImplemented +} +func (s *Server) SelfSetTalkPassword(context.Context, string, string) error { + return ErrNotImplemented +} +func (s *Server) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) { + return "", ErrNotImplemented +} +func (s *Server) SelfLogout(context.Context, string) error { return ErrNotImplemented } +func (s *Server) UnlockTalk(context.Context, string, string, string) error { + return ErrNotImplemented +} +func (s *Server) HasTalkGrant(context.Context, string, string) (bool, error) { + return false, nil +} +func (s *Server) Disable(context.Context, string) error { return ErrNotImplemented } +func (s *Server) Enable(context.Context, string) error { return ErrNotImplemented } +func (s *Server) Delete(context.Context, string) error { return ErrNotImplemented } + +var _ Service = (*Server)(nil) +var _ http.Handler = (*RegisterHandler)(nil) + +func setCORS(w http.ResponseWriter) { + w.Header().Set("Access-Control-Allow-Origin", "*") +} + +func writeRegisterOK(w http.ResponseWriter, result RegisterResult) { + setCORS(w) + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(http.StatusOK) + _ = protocol.Encode(w, protocol.RegisterResponse{ + OK: true, + Data: protocol.RegisterData{ + ID: result.ID, + LoginPassword: result.LoginPassword, + }, + }) +} + +func writeRegisterError(w http.ResponseWriter, err *APIError) { + setCORS(w) + w.Header().Set("Content-Type", "application/json; charset=utf-8") + status := http.StatusInternalServerError + code, msg := protocol.CodeBusy, "internal error" + if err != nil { + status = err.Status + code, msg = err.Code, err.Message + } + w.WriteHeader(status) + _ = protocol.Encode(w, protocol.RegisterResponse{ + OK: false, + Error: &protocol.ErrorBody{Code: code, Message: msg}, + }) +} + +func clientIPFromRemoteAddr(r *http.Request) string { + host, _, err := net.SplitHostPort(r.RemoteAddr) + if err != nil { + return r.RemoteAddr + } + return host +} + +func settingTruthy(v string) bool { + switch strings.TrimSpace(strings.ToLower(v)) { + case "1", "true", "yes", "on": + return true + default: + return false + } +} + +func constantTimeEqual(a, b string) bool { + // 长度不同时 ConstantTimeCompare 直接失败;先按较短侧对齐比较,再核对长度,避免过早返回。 + ab := []byte(a) + bb := []byte(b) + if len(ab) != len(bb) { + dummy := make([]byte, len(ab)) + subtle.ConstantTimeCompare(ab, dummy) + return false + } + return subtle.ConstantTimeCompare(ab, bb) == 1 +} + +func generateEndpointID() (string, error) { + b := make([]byte, 8) + if _, err := rand.Read(b); err != nil { + return "", err + } + out := make([]byte, 8) + for i := range b { + out[i] = idAlphabet[int(b[i])%len(idAlphabet)] + } + return "e_" + string(out), nil +} + +func generateLoginPassword() (string, error) { + // 20 字节可读字符,满足 8–128,且不以 nst_ 开头(字母数字混合,冲突概率极低;若撞前缀则重抽)。 + const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789" + for range 8 { + b := make([]byte, 20) + if _, err := rand.Read(b); err != nil { + return "", err + } + out := make([]byte, 20) + for i := range b { + out[i] = alphabet[int(b[i])%len(alphabet)] + } + pw := string(out) + if !strings.HasPrefix(pw, protocol.SessionTokenPrefix) && utf8.RuneCountInString(pw) >= protocol.MinLoginPasswordLen { + return pw, nil + } + } + return "", errors.New("identity: generate login password failed") +} + +func isUniqueConstraint(err error) bool { + if err == nil { + return false + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "unique constraint") || strings.Contains(msg, "constraint failed") +} + +func isBodyTooLarge(err error) bool { + if err == nil { + return false + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "request body too large") || strings.Contains(msg, "http: request body too large") +} diff --git a/internal/app/identity/register_test.go b/internal/app/identity/register_test.go new file mode 100644 index 0000000..cb49fd6 --- /dev/null +++ b/internal/app/identity/register_test.go @@ -0,0 +1,446 @@ +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) + } +}