package admin import ( "context" "crypto/rand" "database/sql" "errors" "net/http" "strings" "time" "unicode/utf8" "git.asio.asia/nixevol/NixMsg/internal/httpx" ) const ( settingRegistrationEnabled = "registration_enabled" settingRegistrationCode = "registration_code" minRegistrationCodeLen = 8 maxRegistrationCodeLen = 64 generatedRegistrationLen = 16 ) func (h *Handler) handleRegistrationGet(w http.ResponseWriter, r *http.Request) { data, err := h.loadRegistration(r.Context()) if err != nil { httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return } httpx.WriteOK(w, data) } func (h *Handler) handleRegistrationPut(w http.ResponseWriter, r *http.Request) { p, _ := principalFrom(r.Context()) ip := httpx.ClientIP(r, h.trusted) var req struct { Enabled *bool `json:"enabled"` Code *string `json:"code"` Generate bool `json:"generate"` } if err := httpx.DecodeJSON(r, &req); err != nil { h.audit(actorString(p), "registration_update", "", "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") return } if req.Enabled == nil && req.Code == nil && !req.Generate { h.audit(actorString(p), "registration_update", "", "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", "至少提供一项更新") return } nowMs := time.Now().UnixMilli() err := h.db.Queue.Do(r.Context(), func(tx *sql.Tx) error { if req.Enabled != nil { val := "0" if *req.Enabled { val = "1" } if _, e := 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, val, nowMs, ); e != nil { return e } } if req.Generate { code, genErr := generateRegistrationCode() if genErr != nil { return genErr } if _, e := 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, nowMs, ); e != nil { return e } return nil } if req.Code != nil { code := *req.Code n := utf8.RuneCountInString(code) if n < minRegistrationCodeLen || n > maxRegistrationCodeLen { return errBadRequest("安全码须为 8–64 字符") } if _, e := 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, nowMs, ); e != nil { return e } } return nil }) if err != nil { var br badRequestError if errors.As(err, &br) { h.audit(actorString(p), "registration_update", "", "bad_request", ip) httpx.WriteError(w, http.StatusBadRequest, "bad_request", string(br)) return } h.audit(actorString(p), "registration_update", "", "error", ip) httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return } data, err := h.loadRegistration(r.Context()) if err != nil { h.audit(actorString(p), "registration_update", "", "error", ip) httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") return } // 审计不写安全码明文 h.audit(actorString(p), "registration_update", "", "ok", ip) httpx.WriteOK(w, data) } func (h *Handler) loadRegistration(ctx context.Context) (map[string]any, error) { var enabledVal, codeVal sql.NullString var enabledAt, codeAt sql.NullInt64 err := h.db.Read.QueryRowContext(ctx, `SELECT value, updated_at FROM settings WHERE key = ?`, settingRegistrationEnabled, ).Scan(&enabledVal, &enabledAt) if err != nil && !errors.Is(err, sql.ErrNoRows) { return nil, err } err = h.db.Read.QueryRowContext(ctx, `SELECT value, updated_at FROM settings WHERE key = ?`, settingRegistrationCode, ).Scan(&codeVal, &codeAt) if err != nil && !errors.Is(err, sql.ErrNoRows) { return nil, err } updated := int64(0) if enabledAt.Valid && enabledAt.Int64 > updated { updated = enabledAt.Int64 } if codeAt.Valid && codeAt.Int64 > updated { updated = codeAt.Int64 } return map[string]any{ "enabled": registrationTruthy(enabledVal.String), "code": codeVal.String, "updated_at_ms": updated, }, nil } func registrationTruthy(v string) bool { switch strings.TrimSpace(strings.ToLower(v)) { case "1", "true", "yes", "on": return true default: return false } } func generateRegistrationCode() (string, error) { const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789" b := make([]byte, generatedRegistrationLen) if _, err := rand.Read(b); err != nil { return "", err } out := make([]byte, generatedRegistrationLen) for i := range b { out[i] = alphabet[int(b[i])%len(alphabet)] } return string(out), nil }