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, }) }