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

373 lines
11 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/geo"
"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
geo *geo.Geo // IP 地区解析(离线)
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,
geo: geo.New("data/ip2region_v4.xdb", "/opt/speedtest/data/ip2region_v4.xdb"),
}
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
}
// Ping 延迟探测:返回最小响应,供前端计算 RTT;
// 同时回传客户端 IP(页面上方展示用)与其地区
func (h *Handler) Ping(c *gin.Context) {
ip := clientIP(c)
c.Header("Cache-Control", "no-store")
c.JSON(http.StatusOK, gin.H{
"pong": true,
"ts": time.Now().UnixMilli(),
"client_ip": ip,
"location": h.geo.Search(ip),
})
}
// 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,
Location: h.geo.Search(ip),
UserAgent: truncate(c.Request.UserAgent(), 255),
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})
}
// truncate 截断字符串到 n 个字节
func truncate(s string, n int) string {
if len(s) <= n {
return s
}
return s[:n]
}
// parseBrowser 从 User-Agent 解析简短的浏览器 + 系统信息,如 "Chrome · Windows"
func parseBrowser(ua string) string {
lower := strings.ToLower(ua)
browser := "其他"
switch {
case strings.Contains(lower, "edg/"):
browser = "Edge"
case strings.Contains(lower, "chrome"):
browser = "Chrome"
case strings.Contains(lower, "firefox"):
browser = "Firefox"
case strings.Contains(lower, "micromessenger"):
browser = "微信"
case strings.Contains(lower, "qqbrowser") || strings.Contains(lower, " qq/"):
browser = "QQ浏览器"
case strings.Contains(lower, "ucbrowser"):
browser = "UC浏览器"
case strings.Contains(lower, "opera") || strings.Contains(lower, "opr/"):
browser = "Opera"
case strings.Contains(lower, "safari"):
browser = "Safari"
}
osName := "其他系统"
switch {
case strings.Contains(lower, "windows"):
osName = "Windows"
case strings.Contains(lower, "android"):
osName = "Android"
case strings.Contains(lower, "iphone") || strings.Contains(lower, "ipad") || strings.Contains(lower, "ios"):
osName = "iOS"
case strings.Contains(lower, "mac os") || strings.Contains(lower, "macintosh"):
osName = "macOS"
case strings.Contains(lower, "linux"):
osName = "Linux"
}
return browser + " · " + osName
}
// RankItem 排行榜展示条目(不含任何 IP 信息——IP 仅后台保存,不对外展示)
type RankItem struct {
ID uint `json:"id"`
Rank int `json:"rank"` // 名次(含附加记录的真实名次)
IsMine bool `json:"is_mine"` // 是否请求方指定的记录(前端高亮"我的成绩")
Location string `json:"location"` // 用户地区(ip2region
Browser string `json:"browser"` // 浏览器 + 系统(由 UA 解析)
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"`
}
func toRankItems(rows []db.SpeedTestResult) []RankItem {
items := make([]RankItem, 0, len(rows))
for _, r := range rows {
items = append(items, RankItem{
ID: r.ID,
Location: r.Location,
Browser: parseBrowser(r.UserAgent),
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"
case "created":
col = "created_at"
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,
})
}