Files
pangolin/server/internal/admin/handlers.go
T
wangjia f819a77d83 fix(admin): 登录审计记录真实来源 IP(经本机反代时取 XFF 末跳)
管理后台经 caddy mTLS 网关(同机 loopback 反代 127.0.0.1:9444)对外后,
LoginSubmit 里 hostOnly(r.RemoteAddr) 恒为 127.0.0.1,admin_login_ok/fail
的 ip 字段失去溯源价值。

新增 realIP():仅当 TCP 对端是 loopback(请求确实来自本机可信反代)才信
X-Forwarded-For,且取**最后一跳**(Caddy 把真实 TCP 对端追加在末位,更早的
段可被客户端伪造预置);直连、XFF 缺失或非法值一律回退 TCP 对端。

IPAllow 白名单中间件刻意不变——安全闸继续只看 RemoteAddr,不受任何代理头
影响(维持原注释声明的边界)。

测试:TestRealIP 七例(直连/伪造头/单跳/多跳/缺失/非法/IPv6);
go test ./internal/admin 全绿。(internal/store 迁移彩排 3 例失败为 HEAD
既有——主仓干净树复现,与本次无关。)

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01P9G7E3wmAYL9KeYCVZVsqu
2026-07-24 08:34:57 +08:00

465 lines
14 KiB
Go

package admin
import (
"context"
"encoding/csv"
"fmt"
"log"
"net/http"
"net/url"
"strconv"
"strings"
"time"
)
// Handlers holds the dependencies for all admin HTTP handlers.
type Handlers struct {
cfg *Config
store Store
sessions *SessionStore
auth *Authenticator
svc Services
sec *SecurityLog
render *renderer
logger *log.Logger
}
// NewHandlers constructs the handler set.
func NewHandlers(cfg *Config, store Store, sessions *SessionStore, auth *Authenticator,
svc Services, sec *SecurityLog, logger *log.Logger) (*Handlers, error) {
r, err := newRenderer()
if err != nil {
return nil, err
}
if logger == nil {
logger = log.Default()
}
return &Handlers{
cfg: cfg, store: store, sessions: sessions, auth: auth,
svc: svc, sec: sec, render: r, logger: logger,
}, nil
}
const batchListLimit = 50
const auditListLimit = 50
// --------------------------------------------------------------------------
// Auth
// --------------------------------------------------------------------------
// LoginPage renders the login form (GET /login).
func (h *Handlers) LoginPage(w http.ResponseWriter, r *http.Request) {
h.render.render(w, "login", pageData{Flash: r.URL.Query().Get("e")})
}
// LoginSubmit processes the login form (POST /login).
func (h *Handlers) LoginSubmit(w http.ResponseWriter, r *http.Request) {
if err := r.ParseForm(); err != nil {
h.render.render(w, "login", pageData{Flash: "请求无效"})
return
}
username := strings.TrimSpace(r.PostFormValue("username"))
password := r.PostFormValue("password")
code := strings.TrimSpace(r.PostFormValue("totp"))
ip := realIP(r) // 经本机 caddy 反代时取 XFF 末跳,否则 TCP 对端(见 mw_ipallow.go)
sid, _, err := h.auth.Login(r.Context(), username, password, code, ip)
if err != nil {
flash := "用户名、密码或动态验证码有误"
if err == ErrLockedOut {
flash = "尝试过于频繁,账户已临时锁定,请稍后再试"
}
w.WriteHeader(http.StatusUnauthorized)
h.render.render(w, "login", pageData{Flash: flash})
return
}
h.setSessionCookie(w, sid)
http.Redirect(w, r, "/", http.StatusFound)
}
// Logout destroys the session (POST /logout).
func (h *Handlers) Logout(w http.ResponseWriter, r *http.Request) {
sess := SessionFromContext(r.Context())
if sess == nil || !h.validCSRF(r, sess) {
http.Error(w, "invalid csrf", http.StatusForbidden)
return
}
if c, err := r.Cookie(SessionCookieName); err == nil {
_ = h.sessions.Delete(r.Context(), c.Value)
}
h.clearSessionCookie(w)
http.Redirect(w, r, "/login", http.StatusFound)
}
// --------------------------------------------------------------------------
// Dashboard
// --------------------------------------------------------------------------
// Dashboard renders the landing page (GET /).
func (h *Handlers) Dashboard(w http.ResponseWriter, r *http.Request) {
sess := SessionFromContext(r.Context())
h.render.render(w, "dashboard", pageData{
Username: sess.Username, CSRF: sess.CSRFToken, Active: "dashboard",
})
}
// --------------------------------------------------------------------------
// Code batches
// --------------------------------------------------------------------------
type codesView struct {
Batches []BatchSummary
HasPrev bool
HasNext bool
PrevOffset int
NextOffset int
}
// CodesPage lists batches and shows the generate form (GET /codes).
func (h *Handlers) CodesPage(w http.ResponseWriter, r *http.Request) {
sess := SessionFromContext(r.Context())
offset, _ := strconv.Atoi(r.URL.Query().Get("offset"))
if offset < 0 {
offset = 0
}
batches, total, err := h.svc.Codes.ListBatches(r.Context(), batchListLimit, offset)
if err != nil {
h.serverError(w, "list batches", err)
return
}
view := codesView{
Batches: batches,
HasPrev: offset > 0,
PrevOffset: maxInt(0, offset-batchListLimit),
HasNext: offset+batchListLimit < total,
NextOffset: offset + batchListLimit,
}
h.render.render(w, "codes", pageData{
Username: sess.Username, CSRF: sess.CSRFToken, Active: "codes",
Flash: r.URL.Query().Get("flash"), Data: view,
})
}
// CreateBatch generates a batch and streams the plaintext CSV (POST /codes).
// The plaintext codes appear ONLY in this response body — never persisted or
// logged.
func (h *Handlers) CreateBatch(w http.ResponseWriter, r *http.Request) {
sess := SessionFromContext(r.Context())
if err := r.ParseForm(); err != nil {
http.Error(w, "bad request", http.StatusBadRequest)
return
}
if !h.validCSRF(r, sess) {
http.Error(w, "invalid csrf", http.StatusForbidden)
return
}
plan := strings.TrimSpace(r.PostFormValue("plan"))
channel := strings.TrimSpace(r.PostFormValue("channel"))
note := strings.TrimSpace(r.PostFormValue("note"))
duration, derr := strconv.Atoi(r.PostFormValue("duration_days"))
count, cerr := strconv.Atoi(r.PostFormValue("count"))
if derr != nil || cerr != nil || duration < 1 || duration > 3650 || count < 1 || count > 5000 {
http.Error(w, "参数无效:时长 1-3650 天,数量 1-5000", http.StatusBadRequest)
return
}
batch, err := h.svc.Codes.CreateBatch(r.Context(), CodeBatchParams{
Plan: plan, DurationDays: duration, Count: count, Channel: channel,
Note: note, CreatedBy: "admin:" + sess.Username,
})
if err != nil {
h.serverError(w, "create batch", err)
return
}
// Audit: record metadata only (NEVER the plaintext codes).
h.writeAudit(r.Context(), sess.Username, "code_batch_create",
fmt.Sprintf("batch:%d", batch.BatchID),
fmt.Sprintf(`{"plan":%q,"duration_days":%d,"count":%d,"channel":%q}`,
batch.Plan, batch.DurationDays, len(batch.Codes), batch.Channel))
h.streamBatchCSV(w, batch)
}
// streamBatchCSV writes the one-time plaintext CSV download.
func (h *Handlers) streamBatchCSV(w http.ResponseWriter, batch *GeneratedBatch) {
filename := fmt.Sprintf("batch-%d-%s.csv", batch.BatchID, batch.GeneratedAt.Format("20060102T150405Z"))
w.Header().Set("Content-Type", "text/csv; charset=utf-8")
w.Header().Set("Content-Disposition", "attachment; filename=\""+filename+"\"")
w.Header().Set("Cache-Control", "no-store")
cw := csv.NewWriter(w)
_ = cw.Write([]string{"index", "code", "plan", "duration_days", "batch_id", "channel"})
for i, code := range batch.Codes {
_ = cw.Write([]string{
strconv.Itoa(i + 1), code, batch.Plan,
strconv.Itoa(batch.DurationDays), strconv.FormatInt(batch.BatchID, 10), batch.Channel,
})
}
cw.Flush()
}
// VoidBatch voids all unused codes in a batch (POST /codes/void).
func (h *Handlers) VoidBatch(w http.ResponseWriter, r *http.Request) {
sess := SessionFromContext(r.Context())
if !h.checkWrite(w, r, sess) {
return
}
batchID, err := strconv.ParseInt(r.PostFormValue("batch_id"), 10, 64)
if err != nil || batchID <= 0 {
http.Error(w, "bad batch_id", http.StatusBadRequest)
return
}
affected, verr := h.svc.Codes.VoidBatch(r.Context(), batchID)
if verr != nil {
h.serverError(w, "void batch", verr)
return
}
h.writeAudit(r.Context(), sess.Username, "code_batch_void",
fmt.Sprintf("batch:%d", batchID),
fmt.Sprintf(`{"voided":%d}`, affected))
h.redirectFlash(w, r, "/codes", fmt.Sprintf("已作废批次 %d 的 %d 个未使用激活码", batchID, affected))
}
// --------------------------------------------------------------------------
// Nodes
// --------------------------------------------------------------------------
type nodesView struct {
Nodes []NodeRow
ProvisionReady bool
LifecycleReady bool
}
// NodesPage lists nodes and operation controls (GET /nodes).
func (h *Handlers) NodesPage(w http.ResponseWriter, r *http.Request) {
sess := SessionFromContext(r.Context())
nodes, err := h.store.ListNodes(r.Context(), 3)
if err != nil {
h.serverError(w, "list nodes", err)
return
}
h.render.render(w, "nodes", pageData{
Username: sess.Username, CSRF: sess.CSRFToken, Active: "nodes",
Flash: r.URL.Query().Get("flash"),
Data: nodesView{
Nodes: nodes,
ProvisionReady: h.svc.Provision.Ready(),
LifecycleReady: h.svc.Lifecycle.Ready(),
},
})
}
// NodeOp dispatches replace / draining / up (POST /nodes/op).
func (h *Handlers) NodeOp(w http.ResponseWriter, r *http.Request) {
sess := SessionFromContext(r.Context())
if !h.checkWrite(w, r, sess) {
return
}
nodeID, err := strconv.ParseInt(r.PostFormValue("node_id"), 10, 64)
if err != nil || nodeID <= 0 {
http.Error(w, "bad node_id", http.StatusBadRequest)
return
}
op := r.PostFormValue("op")
var opErr error
var action string
switch op {
case "replace":
action = "node_replace"
opErr = h.svc.Provision.Replace(r.Context(), nodeID, sess.Username)
case "draining", "up":
action = "node_" + op
opErr = h.svc.Lifecycle.TransitionStatus(r.Context(), nodeID, op, sess.Username)
default:
http.Error(w, "unknown op", http.StatusBadRequest)
return
}
if opErr != nil {
if opErr == ErrServiceUnavailable {
h.redirectFlash(w, r, "/nodes", "该操作所依赖的服务尚未接入")
return
}
h.serverError(w, "node op", opErr)
return
}
h.writeAudit(r.Context(), sess.Username, action,
fmt.Sprintf("node:%d", nodeID), fmt.Sprintf(`{"op":%q}`, op))
h.redirectFlash(w, r, "/nodes", fmt.Sprintf("节点 %d 操作 %s 已执行", nodeID, op))
}
// --------------------------------------------------------------------------
// Audit
// --------------------------------------------------------------------------
type auditView struct {
Filter AuditFilter
FromStr string
ToStr string
Entries []AuditEntry
Total int
HasPrev bool
HasNext bool
PrevURL string
NextURL string
NodeEvents []NodeEvent
NodeIDStr string
}
// AuditPage renders the filtered audit log and optional node events (GET /audit).
func (h *Handlers) AuditPage(w http.ResponseWriter, r *http.Request) {
sess := SessionFromContext(r.Context())
q := r.URL.Query()
offset, _ := strconv.Atoi(q.Get("offset"))
if offset < 0 {
offset = 0
}
fromStr := strings.TrimSpace(q.Get("from"))
toStr := strings.TrimSpace(q.Get("to"))
f := AuditFilter{
Actor: strings.TrimSpace(q.Get("actor")),
Action: strings.TrimSpace(q.Get("action")),
Target: strings.TrimSpace(q.Get("target")),
From: parseDate(fromStr, false),
To: parseDate(toStr, true),
Limit: auditListLimit,
Offset: offset,
}
entries, total, err := h.store.QueryAudit(r.Context(), f)
if err != nil {
h.serverError(w, "query audit", err)
return
}
var nodeEvents []NodeEvent
nodeIDStr := strings.TrimSpace(q.Get("node_id"))
if nodeIDStr != "" {
if nid, perr := strconv.ParseInt(nodeIDStr, 10, 64); perr == nil && nid > 0 {
nodeEvents, err = h.store.QueryNodeEvents(r.Context(), nid, 50)
if err != nil {
h.serverError(w, "query node events", err)
return
}
}
}
view := auditView{
Filter: f, FromStr: fromStr, ToStr: toStr,
Entries: entries, Total: total,
HasPrev: offset > 0, HasNext: offset+auditListLimit < total,
PrevURL: auditURL(q, maxInt(0, offset-auditListLimit)),
NextURL: auditURL(q, offset+auditListLimit),
NodeEvents: nodeEvents, NodeIDStr: nodeIDStr,
}
h.render.render(w, "audit", pageData{
Username: sess.Username, CSRF: sess.CSRFToken, Active: "audit", Data: view,
})
}
// --------------------------------------------------------------------------
// Helpers
// --------------------------------------------------------------------------
func (h *Handlers) validCSRF(r *http.Request, sess *Session) bool {
return sess != nil && sess.ValidCSRF(r.PostFormValue("csrf_token"))
}
// checkWrite enforces CSRF + explicit second confirmation for write ops.
func (h *Handlers) checkWrite(w http.ResponseWriter, r *http.Request, sess *Session) bool {
if err := r.ParseForm(); err != nil {
http.Error(w, "bad request", http.StatusBadRequest)
return false
}
if !h.validCSRF(r, sess) {
http.Error(w, "invalid csrf", http.StatusForbidden)
return false
}
if r.PostFormValue("confirm") != "yes" {
http.Error(w, "confirmation required", http.StatusBadRequest)
return false
}
return true
}
func (h *Handlers) writeAudit(ctx context.Context, actor, action, target, metaJSON string) {
if err := h.store.WriteAudit(ctx, actor, action, target, metaJSON); err != nil {
h.logger.Printf("admin: audit write failed action=%s target=%s: %v", action, target, err)
}
}
func (h *Handlers) setSessionCookie(w http.ResponseWriter, sid string) {
http.SetCookie(w, &http.Cookie{
Name: SessionCookieName,
Value: sid,
Path: "/",
HttpOnly: true,
Secure: h.cfg.CookieSecure,
SameSite: http.SameSiteStrictMode,
MaxAge: int(h.cfg.SessionTTL.Seconds()),
})
}
func (h *Handlers) clearSessionCookie(w http.ResponseWriter) {
http.SetCookie(w, &http.Cookie{
Name: SessionCookieName,
Value: "",
Path: "/",
HttpOnly: true,
Secure: h.cfg.CookieSecure,
SameSite: http.SameSiteStrictMode,
MaxAge: -1,
})
}
func (h *Handlers) redirectFlash(w http.ResponseWriter, r *http.Request, path, msg string) {
http.Redirect(w, r, path+"?flash="+url.QueryEscape(msg), http.StatusFound)
}
func (h *Handlers) serverError(w http.ResponseWriter, what string, err error) {
h.logger.Printf("admin: %s: %v", what, err)
http.Error(w, "服务器内部错误", http.StatusInternalServerError)
}
func auditURL(q url.Values, offset int) string {
nq := url.Values{}
for _, k := range []string{"actor", "action", "target", "from", "to", "node_id"} {
if v := q.Get(k); v != "" {
nq.Set(k, v)
}
}
nq.Set("offset", strconv.Itoa(offset))
return "/audit?" + nq.Encode()
}
// parseDate accepts "2006-01-02" or RFC3339. When endOfDay is true a bare date
// is pushed to 23:59:59 so the range is inclusive.
func parseDate(s string, endOfDay bool) *time.Time {
if s == "" {
return nil
}
if t, err := time.ParseInLocation("2006-01-02", s, time.UTC); err == nil {
if endOfDay {
t = t.Add(24*time.Hour - time.Second)
}
return &t
}
if t, err := time.Parse(time.RFC3339, s); err == nil {
tu := t.UTC()
return &tu
}
return nil
}
func maxInt(a, b int) int {
if a > b {
return a
}
return b
}