174 lines
4.8 KiB
Go
174 lines
4.8 KiB
Go
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
|
||
}
|