Files
speedtest/internal/web/handlers/speedtest.go
T

348 lines
9.9 KiB
Go
Raw 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 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
}
// 附加"我的记录"并标记 is_mine(无论是否在榜内,前端均高亮)
if idStr := strings.TrimSpace(c.Query("include_id")); idStr != "" {
if id, err := strconv.ParseUint(idStr, 10, 64); err == nil && id > 0 {
inList := false
for i := range items {
if uint64(items[i].ID) == id {
items[i].IsMine = true // 榜内:直接标记
inList = true
}
}
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,
})
}