feat: 实现端自助注册 HTTP 接口与 F23 验收测试
This commit is contained in:
+36
-1
@@ -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
|
||||
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user