348 lines
9.8 KiB
Go
348 lines
9.8 KiB
Go
package handlers
|
||
|
||
import (
|
||
"crypto/rand"
|
||
"fmt"
|
||
"io"
|
||
"net"
|
||
"net/http"
|
||
"strconv"
|
||
"strings"
|
||
"sync/atomic"
|
||
"time"
|
||
|
||
"speedtest/config"
|
||
"speedtest/internal/db"
|
||
"speedtest/internal/store"
|
||
|
||
"github.com/gin-gonic/gin"
|
||
)
|
||
|
||
const (
|
||
poolChunkSize = 1 << 20 // 1 MB
|
||
poolChunkN = 16 // 16 MB 随机数据池
|
||
rankLimit = 10 // 排行榜条数
|
||
)
|
||
|
||
// Handler 测速 HTTP 处理器
|
||
type Handler struct {
|
||
stores *store.Stores
|
||
cfg config.SpeedtestConfig
|
||
pool [][]byte // 预生成随机数据池(下载测速负载)
|
||
poolNext atomic.Int64
|
||
}
|
||
|
||
// NewHandler creates a new speedtest Handler.
|
||
func NewHandler(stores *store.Stores, cfg config.SpeedtestConfig) (*Handler, error) {
|
||
h := &Handler{stores: stores, cfg: cfg}
|
||
if err := h.buildPool(); err != nil {
|
||
return nil, err
|
||
}
|
||
return h, nil
|
||
}
|
||
|
||
// buildPool 预生成随机数据池,避免下载时实时生成拖慢吞吐
|
||
func (h *Handler) buildPool() error {
|
||
h.pool = make([][]byte, poolChunkN)
|
||
buf := make([]byte, poolChunkSize*poolChunkN)
|
||
if _, err := rand.Read(buf); err != nil {
|
||
return fmt.Errorf("生成随机数据池失败: %w", err)
|
||
}
|
||
for i := 0; i < poolChunkN; i++ {
|
||
h.pool[i] = buf[i*poolChunkSize : (i+1)*poolChunkSize]
|
||
}
|
||
return nil
|
||
}
|
||
|
||
// clientIP 获取真实客户端 IP:优先 Caddy 传递的 X-Real-IP / X-Forwarded-For
|
||
func clientIP(c *gin.Context) string {
|
||
if xr := c.GetHeader("X-Real-IP"); xr != "" {
|
||
if ip := net.ParseIP(strings.TrimSpace(xr)); ip != nil {
|
||
return ip.String()
|
||
}
|
||
}
|
||
if xf := c.GetHeader("X-Forwarded-For"); xf != "" {
|
||
first := strings.TrimSpace(strings.Split(xf, ",")[0])
|
||
if ip := net.ParseIP(first); ip != nil {
|
||
return ip.String()
|
||
}
|
||
}
|
||
host, _, err := net.SplitHostPort(c.Request.RemoteAddr)
|
||
if err != nil {
|
||
return c.Request.RemoteAddr
|
||
}
|
||
return host
|
||
}
|
||
|
||
// maskIP 排行榜展示时对 IP 打码(保留前 3 段)
|
||
func maskIP(ip string) string {
|
||
if strings.Contains(ip, ":") {
|
||
// IPv6:保留前 3 段
|
||
parts := strings.Split(ip, ":")
|
||
if len(parts) > 3 {
|
||
return strings.Join(parts[:3], ":") + ":****"
|
||
}
|
||
return ip
|
||
}
|
||
parts := strings.Split(ip, ".")
|
||
if len(parts) == 4 {
|
||
return strings.Join(parts[:3], ".") + ".*"
|
||
}
|
||
return ip
|
||
}
|
||
|
||
// Ping 延迟探测:返回最小响应,供前端计算 RTT;同时回传服务器看到的客户端 IP
|
||
// (NAT 环境下可能是内网 IP,前端会用公网探测结果替代)
|
||
func (h *Handler) Ping(c *gin.Context) {
|
||
c.Header("Cache-Control", "no-store")
|
||
c.JSON(http.StatusOK, gin.H{
|
||
"pong": true,
|
||
"ts": time.Now().UnixMilli(),
|
||
"client_ip": clientIP(c),
|
||
})
|
||
}
|
||
|
||
// parseSize 解析 size 参数(字节),限制在 [1, max]
|
||
func parseSize(raw string, def, max int64) int64 {
|
||
size, err := strconv.ParseInt(raw, 10, 64)
|
||
if err != nil || size <= 0 {
|
||
size = def
|
||
}
|
||
if size > max {
|
||
size = max
|
||
}
|
||
return size
|
||
}
|
||
|
||
// Download 下载测速:流式输出预生成的随机数据
|
||
func (h *Handler) Download(c *gin.Context) {
|
||
max := h.cfg.MaxDownloadBytes
|
||
if max <= 0 {
|
||
max = config.DefaultMaxDownloadBytes
|
||
}
|
||
size := parseSize(c.Query("size"), 10*1024*1024, max)
|
||
|
||
c.Header("Content-Type", "application/octet-stream")
|
||
c.Header("Content-Length", strconv.FormatInt(size, 10))
|
||
c.Header("Cache-Control", "no-store, no-cache, must-revalidate, max-age=0")
|
||
c.Header("Content-Disposition", "attachment; filename=random.dat")
|
||
c.Header("Access-Control-Allow-Origin", "*")
|
||
|
||
// 每次请求从不同偏移开始,避免数据完全重复
|
||
start := h.poolNext.Add(1)
|
||
|
||
w := c.Writer
|
||
written := int64(0)
|
||
for written < size {
|
||
idx := (start + written/poolChunkSize) % int64(poolChunkN)
|
||
chunk := h.pool[idx]
|
||
remaining := size - written
|
||
if remaining < int64(len(chunk)) {
|
||
chunk = chunk[:remaining]
|
||
}
|
||
if n, err := w.Write(chunk); err != nil {
|
||
return // 客户端断开,静默结束
|
||
} else {
|
||
written += int64(n)
|
||
}
|
||
}
|
||
}
|
||
|
||
// Upload 上传测速:接收并丢弃请求体,统计接收字节数与耗时
|
||
func (h *Handler) Upload(c *gin.Context) {
|
||
max := h.cfg.MaxUploadBytes
|
||
if max <= 0 {
|
||
max = config.DefaultMaxUploadBytes
|
||
}
|
||
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, max)
|
||
|
||
start := time.Now()
|
||
n, err := io.Copy(io.Discard, c.Request.Body)
|
||
if err != nil {
|
||
c.JSON(http.StatusBadRequest, gin.H{"error": "读取上传数据失败", "detail": err.Error()})
|
||
return
|
||
}
|
||
elapsed := time.Since(start).Seconds()
|
||
mbps := float64(n) * 8 / elapsed / 1e6
|
||
|
||
c.Header("Access-Control-Allow-Origin", "*")
|
||
c.JSON(http.StatusOK, gin.H{
|
||
"received": n,
|
||
"elapsed_s": elapsed,
|
||
"mbps": mbps,
|
||
})
|
||
}
|
||
|
||
// resultReq 前端提交的测速结果
|
||
type resultReq struct {
|
||
LatencyMs float64 `json:"latency_ms"`
|
||
JitterMs float64 `json:"jitter_ms"`
|
||
DownloadMbps float64 `json:"download_mbps"`
|
||
UploadMbps float64 `json:"upload_mbps"`
|
||
ClientIP string `json:"client_ip"` // 可选:前端探测到的公网出口 IP
|
||
ServerIP string `json:"server_ip"` // 可选:服务器看到的来源 IP(NAT 下为内网 IP)
|
||
}
|
||
|
||
// Result 保存一次测速结果
|
||
func (h *Handler) Result(c *gin.Context) {
|
||
var req resultReq
|
||
if err := c.ShouldBindJSON(&req); err != nil {
|
||
c.JSON(http.StatusBadRequest, gin.H{"error": "参数格式错误"})
|
||
return
|
||
}
|
||
|
||
// 基本合理性校验,防止脏数据
|
||
if req.LatencyMs < 0 || req.LatencyMs > 10000 ||
|
||
req.JitterMs < 0 || req.JitterMs > 10000 ||
|
||
req.DownloadMbps < 0 || req.DownloadMbps > 100000 ||
|
||
req.UploadMbps < 0 || req.UploadMbps > 100000 {
|
||
c.JSON(http.StatusBadRequest, gin.H{"error": "参数超出合理范围"})
|
||
return
|
||
}
|
||
|
||
// 客户端展示 IP:优先采用前端探测的公网出口 IP(NAT 环境下服务器只能看到内网 IP);
|
||
// 未提供或格式非法(含回环地址伪造)时回退到 Caddy 传递的 X-Real-IP
|
||
ip := strings.TrimSpace(req.ClientIP)
|
||
if parsed := net.ParseIP(ip); parsed == nil || parsed.IsLoopback() {
|
||
ip = clientIP(c)
|
||
}
|
||
|
||
// 服务器看到的来源 IP(内网 IP,用于调试区分设备):
|
||
// 未提供时回退到 Caddy 传递的 X-Real-IP
|
||
serverIP := strings.TrimSpace(req.ServerIP)
|
||
if net.ParseIP(serverIP) == nil {
|
||
serverIP = clientIP(c)
|
||
}
|
||
|
||
record := &db.SpeedTestResult{
|
||
ClientIP: ip,
|
||
ServerIP: serverIP,
|
||
LatencyMs: req.LatencyMs,
|
||
JitterMs: req.JitterMs,
|
||
DownloadMbps: req.DownloadMbps,
|
||
UploadMbps: req.UploadMbps,
|
||
}
|
||
if err := h.stores.Results.Create(record); err != nil {
|
||
c.JSON(http.StatusInternalServerError, gin.H{"error": "保存结果失败"})
|
||
return
|
||
}
|
||
|
||
c.JSON(http.StatusOK, gin.H{"ok": true, "id": record.ID})
|
||
}
|
||
|
||
// RankItem 排行榜展示条目
|
||
// ClientIP 打码展示;ServerIP 为来源 IP(内网 IP 完整返回,供站长调试区分设备)
|
||
type RankItem struct {
|
||
ID uint `json:"id"`
|
||
Rank int `json:"rank"` // 名次(含附加记录的真实名次)
|
||
IsMine bool `json:"is_mine"` // 是否请求方指定的记录(前端高亮"我的成绩")
|
||
ClientIP string `json:"client_ip"`
|
||
ServerIP string `json:"server_ip"`
|
||
IsPrivateIP bool `json:"is_private_ip"` // ClientIP 是否为内网(探测失败回退场景)
|
||
LatencyMs float64 `json:"latency_ms"`
|
||
JitterMs float64 `json:"jitter_ms"`
|
||
DownloadMbps float64 `json:"download_mbps"`
|
||
UploadMbps float64 `json:"upload_mbps"`
|
||
CreatedAt string `json:"created_at"`
|
||
}
|
||
|
||
// isPrivateIP 判断 IP 是否为内网/保留地址(NAT 网关、局域网、回环等)
|
||
func isPrivateIP(ip string) bool {
|
||
parsed := net.ParseIP(ip)
|
||
if parsed == nil {
|
||
return true
|
||
}
|
||
if parsed.IsLoopback() || parsed.IsLinkLocalUnicast() || parsed.IsLinkLocalMulticast() {
|
||
return true
|
||
}
|
||
if parsed.IsPrivate() || parsed.IsUnspecified() {
|
||
return true
|
||
}
|
||
// IPv4 兼容段(IsPrivate 已覆盖 10/8、172.16/12、192.168/16)
|
||
return false
|
||
}
|
||
|
||
func toRankItems(rows []db.SpeedTestResult) []RankItem {
|
||
items := make([]RankItem, 0, len(rows))
|
||
for _, r := range rows {
|
||
items = append(items, RankItem{
|
||
ID: r.ID,
|
||
ClientIP: maskIP(r.ClientIP),
|
||
ServerIP: r.ServerIP,
|
||
IsPrivateIP: isPrivateIP(r.ClientIP),
|
||
LatencyMs: r.LatencyMs,
|
||
JitterMs: r.JitterMs,
|
||
DownloadMbps: r.DownloadMbps,
|
||
UploadMbps: r.UploadMbps,
|
||
CreatedAt: r.CreatedAt.Format("2006-01-02 15:04"),
|
||
})
|
||
}
|
||
return items
|
||
}
|
||
|
||
// Rankings 排行榜:单表多指标,按 sort/order 动态排序;
|
||
// include_id 指定"我的记录"——不在榜内时附加到末尾(带真实名次与 is_mine 标记),供前端高亮
|
||
func (h *Handler) Rankings(c *gin.Context) {
|
||
sortField := c.DefaultQuery("sort", "download")
|
||
order := c.DefaultQuery("order", "desc")
|
||
|
||
// 列名白名单,防注入
|
||
var col string
|
||
switch sortField {
|
||
case "upload":
|
||
col = "upload_mbps"
|
||
case "latency":
|
||
col = "latency_ms"
|
||
default:
|
||
sortField, col = "download", "download_mbps"
|
||
}
|
||
asc := order == "asc"
|
||
|
||
rows, _ := h.stores.Results.TopBy(col, asc, rankLimit)
|
||
items := toRankItems(rows)
|
||
for i := range items {
|
||
items[i].Rank = i + 1
|
||
}
|
||
|
||
// 附加"我的记录"(不在榜内时)
|
||
if idStr := strings.TrimSpace(c.Query("include_id")); idStr != "" {
|
||
if id, err := strconv.ParseUint(idStr, 10, 64); err == nil && id > 0 {
|
||
inList := false
|
||
for _, it := range items {
|
||
if uint64(it.ID) == id {
|
||
inList = true
|
||
break
|
||
}
|
||
}
|
||
if !inList {
|
||
if rec, err := h.stores.Results.GetByID(uint(id)); err == nil {
|
||
item := toRankItems([]db.SpeedTestResult{*rec})[0]
|
||
var val float64
|
||
switch col {
|
||
case "upload_mbps":
|
||
val = rec.UploadMbps
|
||
case "latency_ms":
|
||
val = rec.LatencyMs
|
||
default:
|
||
val = rec.DownloadMbps
|
||
}
|
||
if better, err := h.stores.Results.CountBetter(col, asc, val); err == nil {
|
||
item.Rank = int(better) + 1 // 真实名次
|
||
}
|
||
item.IsMine = true
|
||
items = append(items, item)
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
total, _ := h.stores.Results.Count()
|
||
c.Header("Cache-Control", "no-store")
|
||
c.JSON(http.StatusOK, gin.H{
|
||
"total": total,
|
||
"sort": sortField,
|
||
"order": order,
|
||
"items": items,
|
||
})
|
||
}
|