284 lines
7.9 KiB
Go
284 lines
7.9 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
|
||
}
|
||
|
||
// 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)
|
||
}
|
||
|
||
record := &db.SpeedTestResult{
|
||
ClientIP: ip,
|
||
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 排行榜展示条目(IP 打码)
|
||
type RankItem struct {
|
||
ID uint `json:"id"`
|
||
ClientIP string `json:"client_ip"`
|
||
IsPrivateIP bool `json:"is_private_ip"` // NAT/内网环境(如路由器网关),前端标注展示
|
||
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),
|
||
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 排行榜:下载榜 / 上传榜 / 延迟榜 + 总测试次数
|
||
func (h *Handler) Rankings(c *gin.Context) {
|
||
download, _ := h.stores.Results.TopDownload(rankLimit)
|
||
upload, _ := h.stores.Results.TopUpload(rankLimit)
|
||
latency, _ := h.stores.Results.TopLatency(rankLimit)
|
||
total, _ := h.stores.Results.Count()
|
||
|
||
c.Header("Cache-Control", "no-store")
|
||
c.JSON(http.StatusOK, gin.H{
|
||
"total": total,
|
||
"download": toRankItems(download),
|
||
"upload": toRankItems(upload),
|
||
"latency": toRankItems(latency),
|
||
})
|
||
}
|