Files

174 lines
4.8 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}