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") }