mirror of
https://gitee.com/samwaf/SamWaf.git
synced 2026-08-31 01:41:39 +08:00
@@ -496,6 +496,7 @@ func (m *wafSystenService) run() {
|
||||
globalobj.GWAF_RUNTIME_OBJ_WAF_TaskRegistry.RegisterTask(enums.TASK_THREAT_IP_SYNC, waftask.TaskThreatIPSync)
|
||||
globalobj.GWAF_RUNTIME_OBJ_WAF_TaskRegistry.RegisterTask(enums.TASK_ACCESS_CLEAN, waftask.TaskAccessClean)
|
||||
globalobj.GWAF_RUNTIME_OBJ_WAF_TaskRegistry.RegisterTask(enums.TASK_HOSTGUARD_CLEAN_EXPIRED, waftask.TaskHostGuardCleanExpired)
|
||||
globalobj.GWAF_RUNTIME_OBJ_WAF_TaskRegistry.RegisterTask(enums.TASK_TRAFFIC_FLUSH, waftask.TaskTrafficFlush)
|
||||
|
||||
// 进程启动重放:把各启用威胁情报渠道的快照重新灌入系统 ipset(内存态重启会丢) 并重建 WAF 并集
|
||||
go waf_service.WafThreatIPServiceApp.RestoreAllOnStartup()
|
||||
@@ -949,6 +950,12 @@ func (m *wafSystenService) stopSamWaf() {
|
||||
zlog.Warn("App Engine is nil, skipping shutdown")
|
||||
}
|
||||
|
||||
// 站点流量落库:引擎已停、不会再有新字节进来,这里补最后一次,
|
||||
// 把内存里没到 30 秒周期的增量写掉(否则每次重启都会丢掉最后不足一个周期的流量)。
|
||||
zlog.Info("Flush SamWaf Traffic Stats...")
|
||||
waftask.FlushTrafficStats()
|
||||
zlog.Info("Flush SamWaf Traffic Stats finished")
|
||||
|
||||
zlog.Info("Shutdown SamWaf Queue Processors...")
|
||||
// 关闭信号通道,通知所有队列处理协程退出
|
||||
close(global.GWAF_QUEUE_SHUTDOWN_SIGNAL)
|
||||
|
||||
@@ -27,4 +27,5 @@ const (
|
||||
TASK_THREAT_IP_SYNC = "task_threat_ip_sync" //威胁情报IP订阅同步
|
||||
TASK_ACCESS_CLEAN = "task_access_clean" //统一访问认证:清理过期会话/令牌/票据与审计日志
|
||||
TASK_HOSTGUARD_CLEAN_EXPIRED = "task_hostguard_clean_expired" //主机防爆破:解封到期封禁(每分钟,因最短阶梯只有5分钟)
|
||||
TASK_TRAFFIC_FLUSH = "task_traffic_flush" //站点流量计量落库(30秒一次,引擎侧字节计量与日志解耦)
|
||||
)
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
package global
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// 流量计量累加器:与访问日志完全解耦的站点流量归集。
|
||||
//
|
||||
// 为什么不复用日志流做统计(issue #930):
|
||||
// 静态资源、大文件下载、chunked/流式响应在引擎里根本不产生访问日志,
|
||||
// 靠日志字段累加出来的「入站/出站流量」必然只统计到少数 HTML 页,
|
||||
// 表现为「PV 正常涨、流量永远停在 KB 级」。这里改为在引擎侧直接量真实字节。
|
||||
//
|
||||
// 设计要点:
|
||||
// 1. 桶键带 day / hour_time,且**由请求发生时刻决定**,不是落库时刻。
|
||||
// 否则 23:59:50 的流量会在 00:00:10 落库时被算到第二天——总量对、分布错。
|
||||
// 2. 分片加锁:累加只在自己那一片的锁内完成(计数器指针不外泄给锁外使用),
|
||||
// Drain 换走整张 map 后,旧 map 不可能再被写入,因此不会丢字节。
|
||||
// 3. 热路径无 DB、无日志、无分配(桶已存在时)。
|
||||
|
||||
// TrafficKey 一个流量桶的键:站点 + 天 + 整点
|
||||
type TrafficKey struct {
|
||||
HostCode string // 网站唯一码
|
||||
Host string // 域名(落库时写入,便于旧行缺失时新建)
|
||||
Day int // 年月日,如 20260818
|
||||
HourTime int64 // 整点 unix 时间戳(秒)
|
||||
}
|
||||
|
||||
// TrafficSnapshot Drain 出来的一个桶的增量
|
||||
type TrafficSnapshot struct {
|
||||
TrafficKey
|
||||
In int64 // 入站字节
|
||||
Out int64 // 出站字节
|
||||
}
|
||||
|
||||
const trafficShardCount = 32
|
||||
|
||||
type trafficCounter struct {
|
||||
in int64
|
||||
out int64
|
||||
}
|
||||
|
||||
type trafficShard struct {
|
||||
mu sync.Mutex
|
||||
m map[TrafficKey]*trafficCounter
|
||||
}
|
||||
|
||||
var trafficShards [trafficShardCount]trafficShard
|
||||
|
||||
// trafficShardOf 按 host_code 取分片,保证同一站点始终落同一片(减少跨片抖动)
|
||||
func trafficShardOf(hostCode string) *trafficShard {
|
||||
// FNV-1a,避免引入额外依赖
|
||||
var h uint32 = 2166136261
|
||||
for i := 0; i < len(hostCode); i++ {
|
||||
h ^= uint32(hostCode[i])
|
||||
h *= 16777619
|
||||
}
|
||||
return &trafficShards[h&(trafficShardCount-1)]
|
||||
}
|
||||
|
||||
// TrafficBucketOf 由「请求发生时刻」算出天/整点,与 weblog.Day、UNIX_ADD_TIME 同源
|
||||
func TrafficBucketOf(t time.Time) (day int, hourTime int64) {
|
||||
y, m, d := t.Date()
|
||||
return y*10000 + int(m)*100 + d, (t.Unix() / 3600) * 3600
|
||||
}
|
||||
|
||||
// AddTraffic 累加一次请求的进出字节。hostCode 为空(未匹配到站点)直接丢弃,无处归属。
|
||||
func AddTraffic(hostCode, host string, day int, hourTime int64, in, out int64) {
|
||||
if hostCode == "" || (in <= 0 && out <= 0) {
|
||||
return
|
||||
}
|
||||
if in < 0 {
|
||||
in = 0
|
||||
}
|
||||
if out < 0 {
|
||||
out = 0
|
||||
}
|
||||
k := TrafficKey{HostCode: hostCode, Host: host, Day: day, HourTime: hourTime}
|
||||
s := trafficShardOf(hostCode)
|
||||
s.mu.Lock()
|
||||
if s.m == nil {
|
||||
s.m = make(map[TrafficKey]*trafficCounter, 8)
|
||||
}
|
||||
c := s.m[k]
|
||||
if c == nil {
|
||||
c = &trafficCounter{}
|
||||
s.m[k] = c
|
||||
}
|
||||
c.in += in
|
||||
c.out += out
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
// DrainTraffic 取走全部累计增量并清空。落库侧按增量 UPSERT,所以整体取走是安全的。
|
||||
func DrainTraffic() []TrafficSnapshot {
|
||||
var out []TrafficSnapshot
|
||||
for i := range trafficShards {
|
||||
s := &trafficShards[i]
|
||||
s.mu.Lock()
|
||||
old := s.m
|
||||
if len(old) == 0 {
|
||||
s.mu.Unlock()
|
||||
continue
|
||||
}
|
||||
s.m = make(map[TrafficKey]*trafficCounter, len(old))
|
||||
s.mu.Unlock()
|
||||
|
||||
for k, c := range old {
|
||||
out = append(out, TrafficSnapshot{TrafficKey: k, In: c.in, Out: c.out})
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// RestoreTraffic 落库失败时把增量放回累加器,等下个周期重试,避免直接丢数。
|
||||
func RestoreTraffic(list []TrafficSnapshot) {
|
||||
for _, s := range list {
|
||||
AddTraffic(s.HostCode, s.Host, s.Day, s.HourTime, s.In, s.Out)
|
||||
}
|
||||
}
|
||||
|
||||
// PendingTrafficBuckets 当前待落库的桶数量(诊断/测试用)
|
||||
func PendingTrafficBuckets() int {
|
||||
n := 0
|
||||
for i := range trafficShards {
|
||||
s := &trafficShards[i]
|
||||
s.mu.Lock()
|
||||
n += len(s.m)
|
||||
s.mu.Unlock()
|
||||
}
|
||||
return n
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
package global
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// 每个用例开头先清空,避免包内用例互相污染
|
||||
func resetTraffic() { DrainTraffic() }
|
||||
|
||||
// 分桶必须由「请求发生时刻」决定:这是本次改造最容易写错、错了又最难查的地方
|
||||
// (落库时刻分桶会把 23:59:50 的流量记到第二天,总量对、分布错)。
|
||||
func TestTrafficBucketOf(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
ts time.Time
|
||||
wantDay int
|
||||
wantHour time.Time
|
||||
}{
|
||||
{"整点前一秒", time.Date(2026, 8, 18, 23, 59, 59, 0, time.Local), 20260818, time.Date(2026, 8, 18, 23, 0, 0, 0, time.Local)},
|
||||
{"跨天后一秒", time.Date(2026, 8, 19, 0, 0, 1, 0, time.Local), 20260819, time.Date(2026, 8, 19, 0, 0, 0, 0, time.Local)},
|
||||
{"月初", time.Date(2026, 9, 1, 12, 34, 56, 0, time.Local), 20260901, time.Date(2026, 9, 1, 12, 0, 0, 0, time.Local)},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
day, hour := TrafficBucketOf(c.ts)
|
||||
if day != c.wantDay {
|
||||
t.Fatalf("day = %d, 期望 %d", day, c.wantDay)
|
||||
}
|
||||
if hour != c.wantHour.Unix() {
|
||||
t.Fatalf("hourTime = %d(%s), 期望 %d(%s)",
|
||||
hour, time.Unix(hour, 0), c.wantHour.Unix(), c.wantHour)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 同一秒的两次请求必须落同一个桶;跨天/跨整点必须落不同桶
|
||||
func TestAddTraffic_BucketSeparation(t *testing.T) {
|
||||
resetTraffic()
|
||||
d1, h1 := TrafficBucketOf(time.Date(2026, 8, 18, 23, 59, 50, 0, time.Local))
|
||||
d2, h2 := TrafficBucketOf(time.Date(2026, 8, 19, 0, 0, 10, 0, time.Local))
|
||||
|
||||
AddTraffic("h1", "a.com", d1, h1, 100, 200)
|
||||
AddTraffic("h1", "a.com", d1, h1, 1, 2) // 同桶累加
|
||||
AddTraffic("h1", "a.com", d2, h2, 10, 20)
|
||||
|
||||
got := DrainTraffic()
|
||||
if len(got) != 2 {
|
||||
t.Fatalf("期望 2 个桶(跨天必须分开),实际 %d 个: %+v", len(got), got)
|
||||
}
|
||||
for _, s := range got {
|
||||
switch s.Day {
|
||||
case d1:
|
||||
if s.In != 101 || s.Out != 202 {
|
||||
t.Fatalf("旧一天的桶 in/out = %d/%d,期望 101/202", s.In, s.Out)
|
||||
}
|
||||
case d2:
|
||||
if s.In != 10 || s.Out != 20 {
|
||||
t.Fatalf("新一天的桶 in/out = %d/%d,期望 10/20", s.In, s.Out)
|
||||
}
|
||||
default:
|
||||
t.Fatalf("出现意外的 day=%d", s.Day)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 不同站点各记各的账,不能串
|
||||
func TestAddTraffic_PerHost(t *testing.T) {
|
||||
resetTraffic()
|
||||
day, hour := TrafficBucketOf(time.Now())
|
||||
AddTraffic("hostA", "a.com", day, hour, 5, 7)
|
||||
AddTraffic("hostB", "b.com", day, hour, 50, 70)
|
||||
|
||||
got := DrainTraffic()
|
||||
if len(got) != 2 {
|
||||
t.Fatalf("期望 2 个站点桶,实际 %d", len(got))
|
||||
}
|
||||
m := map[string][2]int64{}
|
||||
for _, s := range got {
|
||||
m[s.HostCode] = [2]int64{s.In, s.Out}
|
||||
if s.HostCode == "hostA" && s.Host != "a.com" {
|
||||
t.Fatalf("host 字段串了: %s", s.Host)
|
||||
}
|
||||
}
|
||||
if m["hostA"] != [2]int64{5, 7} || m["hostB"] != [2]int64{50, 70} {
|
||||
t.Fatalf("站点账目不对: %+v", m)
|
||||
}
|
||||
}
|
||||
|
||||
// Drain 必须是「取走」:取完就清零,再取为空
|
||||
func TestDrainTraffic_ClearsAfterDrain(t *testing.T) {
|
||||
resetTraffic()
|
||||
day, hour := TrafficBucketOf(time.Now())
|
||||
AddTraffic("h", "a.com", day, hour, 1, 1)
|
||||
if n := PendingTrafficBuckets(); n != 1 {
|
||||
t.Fatalf("待落库桶数 = %d,期望 1", n)
|
||||
}
|
||||
if got := DrainTraffic(); len(got) != 1 {
|
||||
t.Fatalf("首次 Drain 应拿到 1 个桶,实际 %d", len(got))
|
||||
}
|
||||
if n := PendingTrafficBuckets(); n != 0 {
|
||||
t.Fatalf("Drain 后应清零,实际剩 %d 个桶", n)
|
||||
}
|
||||
if got := DrainTraffic(); len(got) != 0 {
|
||||
t.Fatalf("二次 Drain 应为空,实际 %d 个桶(会造成重复计数)", len(got))
|
||||
}
|
||||
}
|
||||
|
||||
// 无效输入直接丢弃:没有 host_code 的流量无处归属,负数/全零不该建桶
|
||||
func TestAddTraffic_IgnoresInvalid(t *testing.T) {
|
||||
resetTraffic()
|
||||
day, hour := TrafficBucketOf(time.Now())
|
||||
AddTraffic("", "a.com", day, hour, 100, 100) // 无 host_code
|
||||
AddTraffic("h", "a.com", day, hour, 0, 0) // 全零
|
||||
AddTraffic("h", "a.com", day, hour, -5, -5) // 负数
|
||||
if n := PendingTrafficBuckets(); n != 0 {
|
||||
t.Fatalf("无效输入不应建桶,实际 %d 个", n)
|
||||
}
|
||||
|
||||
// 一正一负:负的那侧按 0 计,不能把总量拉低
|
||||
AddTraffic("h", "a.com", day, hour, -5, 100)
|
||||
got := DrainTraffic()
|
||||
if len(got) != 1 || got[0].In != 0 || got[0].Out != 100 {
|
||||
t.Fatalf("期望 in=0 out=100,实际 %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
// 落库失败要能原样退回,等下轮重试
|
||||
func TestRestoreTraffic(t *testing.T) {
|
||||
resetTraffic()
|
||||
day, hour := TrafficBucketOf(time.Now())
|
||||
AddTraffic("h", "a.com", day, hour, 11, 22)
|
||||
drained := DrainTraffic()
|
||||
|
||||
RestoreTraffic(drained)
|
||||
again := DrainTraffic()
|
||||
if len(again) != 1 || again[0].In != 11 || again[0].Out != 22 {
|
||||
t.Fatalf("退回后应能重新取到同样的账,实际 %+v", again)
|
||||
}
|
||||
}
|
||||
|
||||
// 并发累加不能丢字节(热路径是多协程同时写)
|
||||
func TestAddTraffic_Concurrent(t *testing.T) {
|
||||
resetTraffic()
|
||||
day, hour := TrafficBucketOf(time.Now())
|
||||
const goroutines = 32
|
||||
const perG = 500
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for g := 0; g < goroutines; g++ {
|
||||
wg.Add(1)
|
||||
go func(g int) {
|
||||
defer wg.Done()
|
||||
// 一半打同一个站点(最坏争用),一半分散到不同站点
|
||||
hostCode := "shared"
|
||||
if g%2 == 1 {
|
||||
hostCode = "host" + string(rune('A'+g))
|
||||
}
|
||||
for i := 0; i < perG; i++ {
|
||||
AddTraffic(hostCode, "x.com", day, hour, 1, 2)
|
||||
}
|
||||
}(g)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
var totalIn, totalOut int64
|
||||
for _, s := range DrainTraffic() {
|
||||
totalIn += s.In
|
||||
totalOut += s.Out
|
||||
}
|
||||
wantIn := int64(goroutines * perG)
|
||||
if totalIn != wantIn || totalOut != wantIn*2 {
|
||||
t.Fatalf("并发累加丢字节:in=%d(期望 %d) out=%d(期望 %d)", totalIn, wantIn, totalOut, wantIn*2)
|
||||
}
|
||||
}
|
||||
|
||||
// 并发 Add 与 Drain 交错时也不能丢:Drain 换走 map 的瞬间必须与写入互斥
|
||||
func TestAddTraffic_ConcurrentWithDrain(t *testing.T) {
|
||||
resetTraffic()
|
||||
day, hour := TrafficBucketOf(time.Now())
|
||||
const writers = 8
|
||||
const perW = 2000
|
||||
|
||||
var collected int64
|
||||
var mu sync.Mutex
|
||||
stop := make(chan struct{})
|
||||
var drainWg sync.WaitGroup
|
||||
drainWg.Add(1)
|
||||
go func() {
|
||||
defer drainWg.Done()
|
||||
for {
|
||||
select {
|
||||
case <-stop:
|
||||
for _, s := range DrainTraffic() { // 收尾再取一次
|
||||
mu.Lock()
|
||||
collected += s.In
|
||||
mu.Unlock()
|
||||
}
|
||||
return
|
||||
default:
|
||||
for _, s := range DrainTraffic() {
|
||||
mu.Lock()
|
||||
collected += s.In
|
||||
mu.Unlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for w := 0; w < writers; w++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < perW; i++ {
|
||||
AddTraffic("h", "a.com", day, hour, 1, 0)
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
close(stop)
|
||||
drainWg.Wait()
|
||||
|
||||
if want := int64(writers * perW); collected != want {
|
||||
t.Fatalf("Add 与 Drain 交错丢字节:收到 %d,期望 %d", collected, want)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkAddTraffic(b *testing.B) {
|
||||
day, hour := TrafficBucketOf(time.Now())
|
||||
b.RunParallel(func(pb *testing.PB) {
|
||||
for pb.Next() {
|
||||
AddTraffic("bench", "a.com", day, hour, 1024, 4096)
|
||||
}
|
||||
})
|
||||
DrainTraffic()
|
||||
}
|
||||
@@ -543,6 +543,46 @@ func RunTaskInitMigrations(db *gorm.DB) error {
|
||||
return tx.Where("task_method = ?", enums.TASK_ACCESS_CLEAN).Delete(&model.Task{}).Error
|
||||
},
|
||||
},
|
||||
// 迁移: 站点流量计量落库任务
|
||||
// 30 秒一次:内存里累计的真实进出字节按天/小时增量落库。周期越短掉进程时丢得越少,
|
||||
// 但每轮只有 2N 条 UPDATE(N=有流量的站点数),30 秒对 SQLite 毫无压力。
|
||||
{
|
||||
ID: "202608180001_add_traffic_flush_task",
|
||||
Migrate: func(tx *gorm.DB) error {
|
||||
zlog.Info("迁移 202608180001: 创建站点流量落库任务")
|
||||
|
||||
var count int64
|
||||
tx.Model(&model.Task{}).Where("task_method = ?", enums.TASK_TRAFFIC_FLUSH).Count(&count)
|
||||
if count > 0 {
|
||||
zlog.Info("站点流量落库任务已存在,跳过", "task_method", enums.TASK_TRAFFIC_FLUSH)
|
||||
return nil
|
||||
}
|
||||
|
||||
task := model.Task{
|
||||
BaseOrm: baseorm.BaseOrm{
|
||||
Id: uuid.GenUUID(),
|
||||
USER_CODE: global.GWAF_USER_CODE,
|
||||
Tenant_ID: global.GWAF_TENANT_ID,
|
||||
CREATE_TIME: customtype.JsonTime(time.Now()),
|
||||
UPDATE_TIME: customtype.JsonTime(time.Now()),
|
||||
},
|
||||
TaskName: "每30秒把站点流量计量落库",
|
||||
TaskUnit: enums.TASK_SECOND,
|
||||
TaskValue: 30,
|
||||
TaskAt: "",
|
||||
TaskMethod: enums.TASK_TRAFFIC_FLUSH,
|
||||
}
|
||||
if err := tx.Create(&task).Error; err != nil {
|
||||
return fmt.Errorf("创建站点流量落库任务失败: %w", err)
|
||||
}
|
||||
zlog.Info("站点流量落库任务创建成功")
|
||||
return nil
|
||||
},
|
||||
Rollback: func(tx *gorm.DB) error {
|
||||
zlog.Info("回滚 202608180001: 删除站点流量落库任务")
|
||||
return tx.Where("task_method = ?", enums.TASK_TRAFFIC_FLUSH).Delete(&model.Task{}).Error
|
||||
},
|
||||
},
|
||||
// 迁移: 主机防爆破的到期解封任务
|
||||
// 1 分钟一次:阶梯最短一级只有 5 分钟,沿用防火墙那个 5 分钟粒度的话,
|
||||
// 用户会看到"明明写着封5分钟,实际封了10分钟"。
|
||||
|
||||
@@ -0,0 +1,194 @@
|
||||
package wafenginecore
|
||||
|
||||
import (
|
||||
"SamWaf/global"
|
||||
"bufio"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
// 站点流量计量:在引擎最外层直接量「真实写给客户端 / 真实从客户端读到」的字节数,
|
||||
// 与「是否记录访问日志」彻底解耦(issue #930)。
|
||||
//
|
||||
// 为什么不能继续靠访问日志累加:
|
||||
// - 静态资源(图片/CSS/JS/音视频,以及任何带 Accept-Ranges 的响应)在 modifyResponse
|
||||
// 里整条日志都不入队,字节自然全丢;
|
||||
// - chunked / 流式响应的 Content-Length 是 -1,日志字段只能记 0;
|
||||
// - 用户把「日志记录类型」调成「只记异常」后,正常流量的账会整体消失。
|
||||
//
|
||||
// 包装层必须把这些接口透传下去,否则是回归事故:
|
||||
// - http.Hijacker —— wafproxy/reverseproxy.go 里协议升级强制要求,不实现 WebSocket 直接 500
|
||||
// - http.Flusher —— 不实现 SSE / 流式响应会被缓冲住
|
||||
// - io.ReaderFrom —— 保住 sendfile/TransmitFile 快路径(http.ServeFile 静态伺服走这条)
|
||||
// - http.Pusher / http.CloseNotifier —— reverseproxy 会做类型断言
|
||||
// - Unwrap —— http.ResponseController 靠它找底层能力
|
||||
|
||||
// trafficMeter 单次请求的字节账本。
|
||||
// 用 atomic:WebSocket 劫持后由两个 copy 协程并发读写同一条连接。
|
||||
type trafficMeter struct {
|
||||
hostCode string
|
||||
host string
|
||||
day int // 请求发生时刻所属的天,不是落库时刻
|
||||
hourTime int64 // 请求发生时刻所属的整点
|
||||
in atomic.Int64
|
||||
out atomic.Int64
|
||||
}
|
||||
|
||||
func (m *trafficMeter) addIn(n int64) {
|
||||
if n > 0 {
|
||||
m.in.Add(n)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *trafficMeter) addOut(n int64) {
|
||||
if n > 0 {
|
||||
m.out.Add(n)
|
||||
}
|
||||
}
|
||||
|
||||
// settle 结账:把账本清零并交给累加器。
|
||||
// 用 Swap 而不是 Load,保证重复调用不会重复计数(劫持连接的迟到字节也只算一次)。
|
||||
func (m *trafficMeter) settle() {
|
||||
in := m.in.Swap(0)
|
||||
out := m.out.Swap(0)
|
||||
global.AddTraffic(m.hostCode, m.host, m.day, m.hourTime, in, out)
|
||||
}
|
||||
|
||||
// countingResponseWriter 计出站字节
|
||||
type countingResponseWriter struct {
|
||||
http.ResponseWriter
|
||||
m trafficMeter // 值内嵌,每请求少一次堆分配
|
||||
}
|
||||
|
||||
func (c *countingResponseWriter) Write(b []byte) (int, error) {
|
||||
n, err := c.ResponseWriter.Write(b)
|
||||
c.m.addOut(int64(n))
|
||||
return n, err
|
||||
}
|
||||
|
||||
// Unwrap 供 http.ResponseController 找到底层 ResponseWriter(SetReadDeadline 等)
|
||||
func (c *countingResponseWriter) Unwrap() http.ResponseWriter {
|
||||
return c.ResponseWriter
|
||||
}
|
||||
|
||||
func (c *countingResponseWriter) Flush() {
|
||||
if f, ok := c.ResponseWriter.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
}
|
||||
|
||||
func (c *countingResponseWriter) Push(target string, opts *http.PushOptions) error {
|
||||
if p, ok := c.ResponseWriter.(http.Pusher); ok {
|
||||
return p.Push(target, opts)
|
||||
}
|
||||
return http.ErrNotSupported
|
||||
}
|
||||
|
||||
// CloseNotify 透传。底层不支持时返回一个永不触发的通道:
|
||||
// 调用方(reverseproxy)只在请求 ctx 无 Done 通道时才会走到这里,
|
||||
// 而真实服务器请求的 ctx 一定是可取消的,因此实际不会命中。
|
||||
func (c *countingResponseWriter) CloseNotify() <-chan bool {
|
||||
if cn, ok := c.ResponseWriter.(http.CloseNotifier); ok {
|
||||
return cn.CloseNotify()
|
||||
}
|
||||
return make(chan bool, 1)
|
||||
}
|
||||
|
||||
// Hijack 透传,并把劫持后的连接也纳入计量(WebSocket 隧道字节全靠这条)
|
||||
func (c *countingResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||
hj, ok := c.ResponseWriter.(http.Hijacker)
|
||||
if !ok {
|
||||
return nil, nil, http.ErrNotSupported
|
||||
}
|
||||
conn, brw, err := hj.Hijack()
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
cc := &countingConn{Conn: conn, m: &c.m}
|
||||
if brw == nil {
|
||||
return cc, nil, nil
|
||||
}
|
||||
|
||||
// http server 已经预读进 bufio 但业务还没消费的字节:它们是从原始 conn 读出来的,
|
||||
// 这里补计一次入站,之后的读走 countingConn。
|
||||
// 必须用 LimitReader 把旧 Reader 截断到已缓冲长度,否则读空缓冲后会继续
|
||||
// 从未计量的原始 conn 上阻塞读取。
|
||||
var src io.Reader = cc
|
||||
if buffered := brw.Reader.Buffered(); buffered > 0 {
|
||||
c.m.addIn(int64(buffered))
|
||||
src = io.MultiReader(io.LimitReader(brw.Reader, int64(buffered)), cc)
|
||||
}
|
||||
// 写侧:net/http 在 Hijack 时给的是一个全新的空 bufio.Writer,没有待刷字节,
|
||||
// 直接换成写向 countingConn 的新 Writer 是安全的。
|
||||
return cc, bufio.NewReadWriter(bufio.NewReader(src), bufio.NewWriter(cc)), nil
|
||||
}
|
||||
|
||||
// writerOnly 屏蔽 ReadFrom,防止 io.Copy 递归回自己
|
||||
type writerOnly struct{ io.Writer }
|
||||
|
||||
// ReadFrom 委托给底层,保住 sendfile/TransmitFile 快路径;底层不支持才退回逐块拷贝
|
||||
func (c *countingResponseWriter) ReadFrom(src io.Reader) (int64, error) {
|
||||
if rf, ok := c.ResponseWriter.(io.ReaderFrom); ok {
|
||||
n, err := rf.ReadFrom(src)
|
||||
c.m.addOut(n)
|
||||
return n, err
|
||||
}
|
||||
n, err := io.Copy(writerOnly{c.ResponseWriter}, src)
|
||||
c.m.addOut(n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
// countingBody 计入站字节(chunked 上传也天然正确,因为量的是真实读出来的量)
|
||||
type countingBody struct {
|
||||
io.ReadCloser
|
||||
m *trafficMeter
|
||||
}
|
||||
|
||||
func (b *countingBody) Read(p []byte) (int, error) {
|
||||
n, err := b.ReadCloser.Read(p)
|
||||
b.m.addIn(int64(n))
|
||||
return n, err
|
||||
}
|
||||
|
||||
// countingConn 劫持后的连接双向计量
|
||||
type countingConn struct {
|
||||
net.Conn
|
||||
m *trafficMeter
|
||||
}
|
||||
|
||||
func (c *countingConn) Read(p []byte) (int, error) {
|
||||
n, err := c.Conn.Read(p)
|
||||
c.m.addIn(int64(n))
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (c *countingConn) Write(p []byte) (int, error) {
|
||||
n, err := c.Conn.Write(p)
|
||||
c.m.addOut(int64(n))
|
||||
return n, err
|
||||
}
|
||||
|
||||
// NetConn 暴露底层连接,供需要具体连接类型的调用方解包
|
||||
func (c *countingConn) NetConn() net.Conn { return c.Conn }
|
||||
|
||||
// attachTrafficMeter 给一次请求挂上计量,返回包装后的 ResponseWriter 与结账函数。
|
||||
// 分桶时间在这里取一次(= 请求发生时刻),绝不能等到落库时再算。
|
||||
func attachTrafficMeter(w http.ResponseWriter, r *http.Request, hostCode, host string) (http.ResponseWriter, func()) {
|
||||
day, hourTime := global.TrafficBucketOf(time.Now())
|
||||
cw := &countingResponseWriter{
|
||||
ResponseWriter: w,
|
||||
m: trafficMeter{
|
||||
hostCode: hostCode,
|
||||
host: host,
|
||||
day: day,
|
||||
hourTime: hourTime,
|
||||
},
|
||||
}
|
||||
if r != nil && r.Body != nil && r.Body != http.NoBody {
|
||||
r.Body = &countingBody{ReadCloser: r.Body, m: &cw.m}
|
||||
}
|
||||
return cw, cw.m.settle
|
||||
}
|
||||
@@ -0,0 +1,353 @@
|
||||
package wafenginecore
|
||||
|
||||
import (
|
||||
"SamWaf/global"
|
||||
"bufio"
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// 端到端计数:起真实 http 服务器,把 attachTrafficMeter 挂上去,
|
||||
// 断言「记到的字节」与「实际收发的字节」逐一对齐。
|
||||
// 覆盖的正是老实现记不到的那些路径:静态大文件、chunked、流式、被拦截响应、WebSocket。
|
||||
|
||||
// tmServe 起一个挂了计量的测试服务器;done 在 handler 结账后关闭
|
||||
func tmServe(hostCode string, h http.HandlerFunc) (*httptest.Server, <-chan struct{}) {
|
||||
done := make(chan struct{})
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
cw, settle := attachTrafficMeter(w, r, hostCode, r.Host)
|
||||
defer close(done) // 后注册先执行:settle 先跑完再放行断言
|
||||
defer settle()
|
||||
h(cw, r)
|
||||
}))
|
||||
return srv, done
|
||||
}
|
||||
|
||||
// 静态大文件:老实现里这类请求连日志都不记,字节全丢(issue #930 主因)
|
||||
func TestTrafficE2E_StaticFileServeFile(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "big.bin")
|
||||
const size = 5 << 20 // 5MB
|
||||
if err := os.WriteFile(path, bytes.Repeat([]byte("A"), size), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
srv, done := tmServe("static", func(w http.ResponseWriter, r *http.Request) {
|
||||
http.ServeFile(w, r, path)
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
resp, err := srv.Client().Get(srv.URL + "/big.bin")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
n, _ := io.Copy(io.Discard, resp.Body)
|
||||
resp.Body.Close()
|
||||
<-done
|
||||
|
||||
if n != size {
|
||||
t.Fatalf("客户端收到 %d 字节,期望 %d", n, size)
|
||||
}
|
||||
_, out := tmDrain("static")
|
||||
if out != size {
|
||||
t.Fatalf("静态大文件出站计数 = %d,期望 %d(差 %d 字节)", out, size, size-out)
|
||||
}
|
||||
}
|
||||
|
||||
// HEAD 请求没有响应体,出站应为 0(不能把 Content-Length 当成已发送字节)
|
||||
func TestTrafficE2E_HeadHasNoBody(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "f.bin")
|
||||
if err := os.WriteFile(path, bytes.Repeat([]byte("B"), 100000), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
srv, done := tmServe("head", func(w http.ResponseWriter, r *http.Request) {
|
||||
http.ServeFile(w, r, path)
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
req, _ := http.NewRequest("HEAD", srv.URL+"/f.bin", nil)
|
||||
resp, err := srv.Client().Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
io.Copy(io.Discard, resp.Body)
|
||||
resp.Body.Close()
|
||||
<-done
|
||||
|
||||
if _, out := tmDrain("head"); out != 0 {
|
||||
t.Fatalf("HEAD 出站计数 = %d,期望 0", out)
|
||||
}
|
||||
}
|
||||
|
||||
// chunked:Content-Length 是 -1,老实现只能记 0
|
||||
func TestTrafficE2E_ChunkedResponse(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
srv, done := tmServe("chunk", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/plain") // 不设 Content-Length → chunked
|
||||
for i := 0; i < 5; i++ {
|
||||
w.Write(bytes.Repeat([]byte("x"), 1000))
|
||||
w.(http.Flusher).Flush()
|
||||
}
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
resp, err := srv.Client().Get(srv.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if resp.ContentLength != -1 {
|
||||
t.Fatalf("前置假设失败:期望 chunked(ContentLength=-1),实际 %d", resp.ContentLength)
|
||||
}
|
||||
n, _ := io.Copy(io.Discard, resp.Body)
|
||||
resp.Body.Close()
|
||||
<-done
|
||||
|
||||
_, out := tmDrain("chunk")
|
||||
if n != 5000 || out != 5000 {
|
||||
t.Fatalf("chunked:客户端收到 %d,计数 %d,期望都是 5000", n, out)
|
||||
}
|
||||
}
|
||||
|
||||
// SSE 流式:每次 Flush 都要能算进去,且 Flush 必须真的透传(否则客户端收不到)
|
||||
func TestTrafficE2E_SSEStreaming(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
const events = 10
|
||||
const oneEvent = "data: 00\n\n"
|
||||
srv, done := tmServe("sse", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
f, ok := w.(http.Flusher)
|
||||
if !ok {
|
||||
t.Error("SSE 场景拿不到 Flusher")
|
||||
return
|
||||
}
|
||||
for i := 0; i < events; i++ {
|
||||
fmt.Fprintf(w, "data: %02d\n\n", i)
|
||||
f.Flush()
|
||||
}
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
resp, err := srv.Client().Get(srv.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 逐事件读,验证确实是流式推送而不是一次性缓冲
|
||||
br := bufio.NewReader(resp.Body)
|
||||
got := 0
|
||||
for i := 0; i < events; i++ {
|
||||
line, err := br.ReadString(byte('\n'))
|
||||
if err != nil {
|
||||
t.Fatalf("第 %d 个事件读取失败: %v", i, err)
|
||||
}
|
||||
if strings.HasPrefix(line, "data: ") {
|
||||
got++
|
||||
}
|
||||
br.ReadString(byte('\n')) // 事件之间的空行
|
||||
}
|
||||
io.Copy(io.Discard, br)
|
||||
resp.Body.Close()
|
||||
<-done
|
||||
|
||||
if got != events {
|
||||
t.Fatalf("只收到 %d 个事件,期望 %d", got, events)
|
||||
}
|
||||
want := int64(events * len(oneEvent))
|
||||
if _, out := tmDrain("sse"); out != want {
|
||||
t.Fatalf("SSE 出站计数 = %d,期望 %d", out, want)
|
||||
}
|
||||
}
|
||||
|
||||
// 被 WAF 拦截的响应(403 + 拦截页)同样要计入出站
|
||||
func TestTrafficE2E_BlockedResponseCounted(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
page := []byte("<html><body>403 blocked by SamWaf</body></html>")
|
||||
srv, done := tmServe("blocked", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(403)
|
||||
w.Write(page)
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
resp, err := srv.Client().Get(srv.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
io.Copy(io.Discard, resp.Body)
|
||||
resp.Body.Close()
|
||||
<-done
|
||||
|
||||
if resp.StatusCode != 403 {
|
||||
t.Fatalf("状态码 = %d", resp.StatusCode)
|
||||
}
|
||||
if _, out := tmDrain("blocked"); out != int64(len(page)) {
|
||||
t.Fatalf("拦截页出站计数 = %d,期望 %d", out, len(page))
|
||||
}
|
||||
}
|
||||
|
||||
// 上传:入站按「真实读到的字节」计
|
||||
func TestTrafficE2E_UploadBody(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
const size = 1 << 20
|
||||
srv, done := tmServe("upload", func(w http.ResponseWriter, r *http.Request) {
|
||||
io.Copy(io.Discard, r.Body)
|
||||
w.WriteHeader(204)
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
resp, err := srv.Client().Post(srv.URL, "application/octet-stream", bytes.NewReader(make([]byte, size)))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
io.Copy(io.Discard, resp.Body)
|
||||
resp.Body.Close()
|
||||
<-done
|
||||
|
||||
in, _ := tmDrain("upload")
|
||||
if in != size {
|
||||
t.Fatalf("上传入站计数 = %d,期望 %d", in, size)
|
||||
}
|
||||
}
|
||||
|
||||
// chunked 上传(Content-Length 未知):老实现取 r.ContentLength = -1 会把入站流量算成负数
|
||||
func TestTrafficE2E_ChunkedUpload(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
const size = 300000
|
||||
srv, done := tmServe("chunkup", func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.ContentLength != -1 {
|
||||
t.Errorf("前置假设失败:期望 chunked 上传 ContentLength=-1,实际 %d", r.ContentLength)
|
||||
}
|
||||
io.Copy(io.Discard, r.Body)
|
||||
w.WriteHeader(204)
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
pr, pw := io.Pipe()
|
||||
go func() {
|
||||
pw.Write(make([]byte, size))
|
||||
pw.Close()
|
||||
}()
|
||||
req, _ := http.NewRequest("POST", srv.URL, pr) // body 长度未知 → chunked
|
||||
resp, err := srv.Client().Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
io.Copy(io.Discard, resp.Body)
|
||||
resp.Body.Close()
|
||||
<-done
|
||||
|
||||
in, _ := tmDrain("chunkup")
|
||||
if in != size {
|
||||
t.Fatalf("chunked 上传入站计数 = %d,期望 %d(且必须为正数)", in, size)
|
||||
}
|
||||
}
|
||||
|
||||
// 口径边界(有意为之,用例锁住行为):只统计业务真正读到的请求体。
|
||||
// 请求被提前拒绝、body 没被读走时,这部分字节不计入入站。
|
||||
func TestTrafficE2E_UnreadBodyNotCounted(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
srv, done := tmServe("unread", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(403) // 直接拒绝,不读 body
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
resp, err := srv.Client().Post(srv.URL, "application/octet-stream", bytes.NewReader(make([]byte, 65536)))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
io.Copy(io.Discard, resp.Body)
|
||||
resp.Body.Close()
|
||||
<-done
|
||||
|
||||
in, _ := tmDrain("unread")
|
||||
if in != 0 {
|
||||
t.Fatalf("未读取的请求体不计入入站,实际计到 %d", in)
|
||||
}
|
||||
}
|
||||
|
||||
// WebSocket 之类的协议升级:劫持之后的隧道字节必须双向计数,且连接不能被包装层弄坏
|
||||
func TestTrafficE2E_HijackedConnectionCounted(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
|
||||
const respHead = "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n"
|
||||
const serverPush = "server-frame-0123456789"
|
||||
const clientPush = "client-frame-abcdef"
|
||||
|
||||
srv, done := tmServe("ws", func(w http.ResponseWriter, r *http.Request) {
|
||||
hj, ok := w.(http.Hijacker)
|
||||
if !ok {
|
||||
t.Error("协议升级拿不到 Hijacker —— WebSocket 会直接 500")
|
||||
return
|
||||
}
|
||||
conn, brw, err := hj.Hijack()
|
||||
if err != nil {
|
||||
t.Errorf("Hijack 失败: %v", err)
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
brw.WriteString(respHead)
|
||||
brw.WriteString(serverPush)
|
||||
if err := brw.Flush(); err != nil {
|
||||
t.Errorf("劫持后写出失败: %v", err)
|
||||
return
|
||||
}
|
||||
buf := make([]byte, len(clientPush))
|
||||
if _, err := io.ReadFull(brw, buf); err != nil {
|
||||
t.Errorf("劫持后读取失败: %v", err)
|
||||
return
|
||||
}
|
||||
if string(buf) != clientPush {
|
||||
t.Errorf("劫持后读到的内容错乱: %q", string(buf))
|
||||
}
|
||||
})
|
||||
defer srv.Close()
|
||||
|
||||
addr := strings.TrimPrefix(srv.URL, "http://")
|
||||
conn, err := net.Dial("tcp", addr)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer conn.Close()
|
||||
conn.SetDeadline(time.Now().Add(10 * time.Second))
|
||||
|
||||
fmt.Fprintf(conn, "GET /ws HTTP/1.1\r\nHost: %s\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n", addr)
|
||||
head := make([]byte, len(respHead)+len(serverPush))
|
||||
if _, err := io.ReadFull(conn, head); err != nil {
|
||||
t.Fatalf("读升级响应失败: %v", err)
|
||||
}
|
||||
if !strings.HasPrefix(string(head), "HTTP/1.1 101") {
|
||||
t.Fatalf("升级响应不对: %q", string(head))
|
||||
}
|
||||
if _, err := conn.Write([]byte(clientPush)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
<-done
|
||||
|
||||
in, out := tmDrain("ws")
|
||||
wantOut := int64(len(respHead) + len(serverPush))
|
||||
if out != wantOut {
|
||||
t.Fatalf("劫持后出站计数 = %d,期望 %d", out, wantOut)
|
||||
}
|
||||
if in < int64(len(clientPush)) {
|
||||
t.Fatalf("劫持后入站计数 = %d,至少应有 %d(客户端推上来的帧)", in, len(clientPush))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
package wafenginecore
|
||||
|
||||
import (
|
||||
"SamWaf/global"
|
||||
"bufio"
|
||||
"bytes"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// ============ 测试脚手架 ============
|
||||
|
||||
// tmDrain 取出某站点这一轮记到的进出字节
|
||||
func tmDrain(hostCode string) (in, out int64) {
|
||||
for _, s := range global.DrainTraffic() {
|
||||
if s.HostCode == hostCode {
|
||||
in += s.In
|
||||
out += s.Out
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// tmBaseRW 最小 ResponseWriter:只记录写了多少
|
||||
type tmBaseRW struct {
|
||||
h http.Header
|
||||
body bytes.Buffer
|
||||
code int
|
||||
}
|
||||
|
||||
func (w *tmBaseRW) Header() http.Header { return w.h }
|
||||
func (w *tmBaseRW) WriteHeader(c int) { w.code = c }
|
||||
func (w *tmBaseRW) Write(b []byte) (int, error) {
|
||||
return w.body.Write(b)
|
||||
}
|
||||
|
||||
// tmFullRW 带 ReaderFrom / Flusher / Pusher / CloseNotifier / Hijacker 的底层
|
||||
type tmFullRW struct {
|
||||
tmBaseRW
|
||||
readFromCalled bool
|
||||
flushCalled bool
|
||||
pushCalled bool
|
||||
closeCh chan bool
|
||||
hijackConn net.Conn
|
||||
hijackBuffered string
|
||||
}
|
||||
|
||||
func (w *tmFullRW) ReadFrom(r io.Reader) (int64, error) {
|
||||
w.readFromCalled = true
|
||||
return io.Copy(&w.body, r)
|
||||
}
|
||||
func (w *tmFullRW) Flush() { w.flushCalled = true }
|
||||
func (w *tmFullRW) CloseNotify() <-chan bool { return w.closeCh }
|
||||
func (w *tmFullRW) Push(string, *http.PushOptions) error {
|
||||
w.pushCalled = true
|
||||
return nil
|
||||
}
|
||||
func (w *tmFullRW) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||
if w.hijackConn == nil {
|
||||
return nil, nil, fmt.Errorf("no conn")
|
||||
}
|
||||
br := bufio.NewReader(io.MultiReader(strings.NewReader(w.hijackBuffered), w.hijackConn))
|
||||
if w.hijackBuffered != "" {
|
||||
_, _ = br.Peek(len(w.hijackBuffered)) // 让缓冲里真的有数据
|
||||
}
|
||||
return w.hijackConn, bufio.NewReadWriter(br, bufio.NewWriter(w.hijackConn)), nil
|
||||
}
|
||||
|
||||
// ============ T10 包装层基本行为 ============
|
||||
|
||||
func TestCountingRW_CountsWrites(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
base := &tmBaseRW{h: http.Header{}}
|
||||
w, settle := attachTrafficMeter(base, httptest.NewRequest("GET", "/", nil), "h1", "a.com")
|
||||
|
||||
w.WriteHeader(200)
|
||||
for i := 0; i < 3; i++ {
|
||||
if _, err := w.Write(make([]byte, 1000)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
settle()
|
||||
|
||||
in, out := tmDrain("h1")
|
||||
if out != 3000 {
|
||||
t.Fatalf("出站计数 = %d,期望 3000", out)
|
||||
}
|
||||
if in != 0 {
|
||||
t.Fatalf("GET 无请求体,入站应为 0,实际 %d", in)
|
||||
}
|
||||
if base.body.Len() != 3000 {
|
||||
t.Fatalf("底层实际写入 %d 字节,包装层不该改变写出内容", base.body.Len())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountingRW_InterfacePassthrough(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
base := &tmFullRW{tmBaseRW: tmBaseRW{h: http.Header{}}, closeCh: make(chan bool, 1)}
|
||||
w, _ := attachTrafficMeter(base, httptest.NewRequest("GET", "/", nil), "h1", "a.com")
|
||||
|
||||
// Flusher:SSE / 流式响应靠它,不透传会被缓冲住
|
||||
f, ok := w.(http.Flusher)
|
||||
if !ok {
|
||||
t.Fatal("包装层必须实现 http.Flusher")
|
||||
}
|
||||
f.Flush()
|
||||
if !base.flushCalled {
|
||||
t.Fatal("Flush 没有透传到底层")
|
||||
}
|
||||
|
||||
// Pusher
|
||||
p, ok := w.(http.Pusher)
|
||||
if !ok {
|
||||
t.Fatal("包装层必须实现 http.Pusher")
|
||||
}
|
||||
if err := p.Push("/x", nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !base.pushCalled {
|
||||
t.Fatal("Push 没有透传到底层")
|
||||
}
|
||||
|
||||
// CloseNotifier:wafproxy/reverseproxy.go 会做类型断言
|
||||
cn, ok := w.(http.CloseNotifier)
|
||||
if !ok {
|
||||
t.Fatal("包装层必须实现 http.CloseNotifier")
|
||||
}
|
||||
if cn.CloseNotify() == nil {
|
||||
t.Fatal("CloseNotify 返回了 nil 通道")
|
||||
}
|
||||
|
||||
// Unwrap:http.ResponseController 靠它找底层能力
|
||||
u, ok := w.(interface{ Unwrap() http.ResponseWriter })
|
||||
if !ok {
|
||||
t.Fatal("包装层必须提供 Unwrap")
|
||||
}
|
||||
if u.Unwrap() != http.ResponseWriter(base) {
|
||||
t.Fatal("Unwrap 没有返回底层 ResponseWriter")
|
||||
}
|
||||
}
|
||||
|
||||
// 底层不支持 Hijack 时必须回 ErrNotSupported,而不是 panic 或假装成功
|
||||
func TestCountingRW_HijackUnsupported(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
base := &tmBaseRW{h: http.Header{}}
|
||||
w, _ := attachTrafficMeter(base, httptest.NewRequest("GET", "/", nil), "h1", "a.com")
|
||||
hj, ok := w.(http.Hijacker)
|
||||
if !ok {
|
||||
t.Fatal("包装层必须实现 http.Hijacker(否则 WebSocket 直接 500)")
|
||||
}
|
||||
if _, _, err := hj.Hijack(); err != http.ErrNotSupported {
|
||||
t.Fatalf("底层不支持时应返回 http.ErrNotSupported,实际 %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ============ T11 ReadFrom 必须委托给底层(保住 sendfile 快路径) ============
|
||||
|
||||
func TestCountingRW_ReadFromDelegates(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
base := &tmFullRW{tmBaseRW: tmBaseRW{h: http.Header{}}}
|
||||
w, settle := attachTrafficMeter(base, httptest.NewRequest("GET", "/", nil), "h1", "a.com")
|
||||
|
||||
rf, ok := w.(io.ReaderFrom)
|
||||
if !ok {
|
||||
t.Fatal("包装层必须实现 io.ReaderFrom")
|
||||
}
|
||||
src := bytes.NewReader(make([]byte, 12345))
|
||||
n, err := rf.ReadFrom(src)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
settle()
|
||||
|
||||
if !base.readFromCalled {
|
||||
t.Fatal("ReadFrom 没有委托给底层:sendfile/TransmitFile 快路径丢了")
|
||||
}
|
||||
if n != 12345 {
|
||||
t.Fatalf("ReadFrom 返回 %d,期望 12345", n)
|
||||
}
|
||||
if _, out := tmDrain("h1"); out != 12345 {
|
||||
t.Fatalf("出站计数 = %d,期望 12345", out)
|
||||
}
|
||||
}
|
||||
|
||||
// 底层没有 ReaderFrom 时要退回逐块拷贝,且计数依然准确
|
||||
func TestCountingRW_ReadFromFallback(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
base := &tmBaseRW{h: http.Header{}}
|
||||
w, settle := attachTrafficMeter(base, httptest.NewRequest("GET", "/", nil), "h1", "a.com")
|
||||
|
||||
n, err := w.(io.ReaderFrom).ReadFrom(bytes.NewReader(make([]byte, 999)))
|
||||
if err != nil || n != 999 {
|
||||
t.Fatalf("ReadFrom = %d, %v", n, err)
|
||||
}
|
||||
settle()
|
||||
if _, out := tmDrain("h1"); out != 999 {
|
||||
t.Fatalf("退回路径出站计数 = %d,期望 999", out)
|
||||
}
|
||||
if base.body.Len() != 999 {
|
||||
t.Fatalf("底层实际收到 %d 字节", base.body.Len())
|
||||
}
|
||||
}
|
||||
|
||||
// ============ T13 结账语义 ============
|
||||
|
||||
// 重复结账不能重复计数(劫持连接可能有迟到字节)
|
||||
func TestTrafficMeter_SettleIsIdempotent(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
base := &tmBaseRW{h: http.Header{}}
|
||||
w, settle := attachTrafficMeter(base, httptest.NewRequest("GET", "/", nil), "h1", "a.com")
|
||||
w.Write(make([]byte, 100))
|
||||
settle()
|
||||
settle()
|
||||
settle()
|
||||
if _, out := tmDrain("h1"); out != 100 {
|
||||
t.Fatalf("重复结账造成重复计数:out = %d,期望 100", out)
|
||||
}
|
||||
}
|
||||
|
||||
// handler panic 时,defer 结账仍然要把已发生的字节记上
|
||||
func TestTrafficMeter_SettleAfterPanic(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
base := &tmBaseRW{h: http.Header{}}
|
||||
func() {
|
||||
defer func() { recover() }()
|
||||
w, settle := attachTrafficMeter(base, httptest.NewRequest("GET", "/", nil), "h1", "a.com")
|
||||
defer settle()
|
||||
w.Write(make([]byte, 512))
|
||||
panic("boom")
|
||||
}()
|
||||
if _, out := tmDrain("h1"); out != 512 {
|
||||
t.Fatalf("panic 后应仍结账 512 字节,实际 %d", out)
|
||||
}
|
||||
}
|
||||
|
||||
// 没有解析到站点(host_code 为空)的流量无处归属,直接丢弃不建桶
|
||||
func TestTrafficMeter_NoHostCodeDropped(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
base := &tmBaseRW{h: http.Header{}}
|
||||
w, settle := attachTrafficMeter(base, httptest.NewRequest("GET", "/", nil), "", "")
|
||||
w.Write(make([]byte, 4096))
|
||||
settle()
|
||||
if n := global.PendingTrafficBuckets(); n != 0 {
|
||||
t.Fatalf("无 host_code 不应建桶,实际 %d 个", n)
|
||||
}
|
||||
}
|
||||
@@ -276,6 +276,13 @@ func (waf *WafEngine) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
incrementMonitor(hostCode)
|
||||
defer decrementMonitor(hostCode)
|
||||
|
||||
// 站点流量计量:直接量真实进出字节,静态资源/大文件/chunked/流式/被拦截响应全都算得到。
|
||||
// 必须在这里(host 已解析、任何响应写出之前)包一次,之后所有出口都用包装后的 w。
|
||||
// 分桶时间取「此刻」,不是落库时刻,否则跨天/跨整点的账会记到错误的桶里。(issue #930)
|
||||
meteredWriter, settleTraffic := attachTrafficMeter(w, r, hostCode, r.Host)
|
||||
w = meteredWriter
|
||||
defer settleTraffic()
|
||||
//检测网站是否已关闭
|
||||
if hostTarget.Host.START_STATUS == 1 {
|
||||
resBytes := []byte("<html><head><title>网站已关闭</title></head><body><center><h1>当前访问网站已关闭</h1> <br><h3></h3></center></body> </html>")
|
||||
@@ -328,8 +335,8 @@ func (waf *WafEngine) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
zlog.Debug("解析cache json失败")
|
||||
}
|
||||
|
||||
// 获取请求报文的内容长度
|
||||
contentLength := r.ContentLength
|
||||
// 获取请求报文的内容长度(chunked 时 Go 给 -1,先归一,避免统计出现负数)
|
||||
contentLength := sanitizeContentLength(r.ContentLength)
|
||||
var bodyByte []byte
|
||||
|
||||
// 拷贝一份request的Body ,控制不记录大文件的情况
|
||||
@@ -883,8 +890,8 @@ func (waf *WafEngine) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
zlog.Debug("write fail:", zap.Any("", err))
|
||||
return
|
||||
}
|
||||
// 获取请求报文的内容长度
|
||||
contentLength := r.ContentLength
|
||||
// 获取请求报文的内容长度(chunked 时 Go 给 -1,先归一,避免统计出现负数)
|
||||
contentLength := sanitizeContentLength(r.ContentLength)
|
||||
|
||||
//server_online[8081].Svr.Close()
|
||||
var bodyByte []byte
|
||||
|
||||
@@ -82,12 +82,13 @@ func CollectStatsFromLogs(logs []*innerbean.WebLog) {
|
||||
Host string
|
||||
}
|
||||
// 站点聚合值
|
||||
// 注意:这里**不再统计流量字节**。日志里静态资源/大文件/chunked/流式响应根本没有记录,
|
||||
// 靠日志累加出来的流量必然偏小(issue #930)。流量改由引擎侧字节计量 + FlushTrafficStats 落库,
|
||||
// 两边写同一张表的不同列,别再把 traffic_* 加回来,否则会双计。
|
||||
type siteAggVal struct {
|
||||
TotalCount int64
|
||||
AttackCount int64
|
||||
NormalCount int64
|
||||
TrafficIn int64
|
||||
TrafficOut int64
|
||||
TotalTimeSpent int64
|
||||
}
|
||||
|
||||
@@ -161,8 +162,6 @@ func CollectStatsFromLogs(logs []*innerbean.WebLog) {
|
||||
} else {
|
||||
sdv.NormalCount++
|
||||
}
|
||||
sdv.TrafficIn += lg.CONTENT_LENGTH
|
||||
sdv.TrafficOut += lg.RES_CONTENT_LENGTH
|
||||
sdv.TotalTimeSpent += lg.TimeSpent
|
||||
|
||||
// 站点小时级聚合(将时间戳截断到整点,注意UNIX_ADD_TIME是毫秒)
|
||||
@@ -184,8 +183,6 @@ func CollectStatsFromLogs(logs []*innerbean.WebLog) {
|
||||
} else {
|
||||
shv.NormalCount++
|
||||
}
|
||||
shv.TrafficIn += lg.CONTENT_LENGTH
|
||||
shv.TrafficOut += lg.RES_CONTENT_LENGTH
|
||||
shv.TotalTimeSpent += lg.TimeSpent
|
||||
}
|
||||
|
||||
@@ -381,8 +378,6 @@ func CollectStatsFromLogs(logs []*innerbean.WebLog) {
|
||||
"total_count": gorm.Expr("total_count + ?", v.TotalCount),
|
||||
"attack_count": gorm.Expr("attack_count + ?", v.AttackCount),
|
||||
"normal_count": gorm.Expr("normal_count + ?", v.NormalCount),
|
||||
"traffic_in": gorm.Expr("traffic_in + ?", v.TrafficIn),
|
||||
"traffic_out": gorm.Expr("traffic_out + ?", v.TrafficOut),
|
||||
"total_time_spent": gorm.Expr("total_time_spent + ?", v.TotalTimeSpent),
|
||||
"update_time": now,
|
||||
})
|
||||
@@ -405,8 +400,6 @@ func CollectStatsFromLogs(logs []*innerbean.WebLog) {
|
||||
TotalCount: v.TotalCount,
|
||||
AttackCount: v.AttackCount,
|
||||
NormalCount: v.NormalCount,
|
||||
TrafficIn: v.TrafficIn,
|
||||
TrafficOut: v.TrafficOut,
|
||||
TotalTimeSpent: v.TotalTimeSpent,
|
||||
}).Error
|
||||
if err != nil {
|
||||
@@ -434,8 +427,6 @@ func CollectStatsFromLogs(logs []*innerbean.WebLog) {
|
||||
"total_count": gorm.Expr("total_count + ?", v.TotalCount),
|
||||
"attack_count": gorm.Expr("attack_count + ?", v.AttackCount),
|
||||
"normal_count": gorm.Expr("normal_count + ?", v.NormalCount),
|
||||
"traffic_in": gorm.Expr("traffic_in + ?", v.TrafficIn),
|
||||
"traffic_out": gorm.Expr("traffic_out + ?", v.TrafficOut),
|
||||
"total_time_spent": gorm.Expr("total_time_spent + ?", v.TotalTimeSpent),
|
||||
"update_time": now,
|
||||
})
|
||||
@@ -458,8 +449,6 @@ func CollectStatsFromLogs(logs []*innerbean.WebLog) {
|
||||
TotalCount: v.TotalCount,
|
||||
AttackCount: v.AttackCount,
|
||||
NormalCount: v.NormalCount,
|
||||
TrafficIn: v.TrafficIn,
|
||||
TrafficOut: v.TrafficOut,
|
||||
TotalTimeSpent: v.TotalTimeSpent,
|
||||
}).Error
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,90 @@
|
||||
package waftask
|
||||
|
||||
import (
|
||||
"go/ast"
|
||||
"go/parser"
|
||||
"go/token"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// 护栏:日志聚合(CollectStatsFromLogs)绝不能再写 traffic_in / traffic_out。
|
||||
//
|
||||
// 背景:流量已改由引擎侧字节计量 + FlushTrafficStats 落库(issue #930)。
|
||||
// 如果哪天有人"顺手"把 traffic 累加加回日志聚合里,两条路径会同时写同一列,
|
||||
// 用户看到的流量会凭空翻倍,而且因为两边都"看起来是对的",极难定位。
|
||||
// 这里用 AST 检查(不看注释,避免误伤说明文字)把这条约束钉死。
|
||||
func TestStatCollectorNoLongerWritesTraffic(t *testing.T) {
|
||||
const file = "stat_collector.go"
|
||||
|
||||
fset := token.NewFileSet()
|
||||
f, err := parser.ParseFile(fset, file, nil, 0) // 不带 ParseComments:注释里提到 traffic 不算违规
|
||||
if err != nil {
|
||||
t.Fatalf("解析 %s 失败: %v", file, err)
|
||||
}
|
||||
|
||||
banned := []string{"traffic_in", "traffic_out", "TrafficIn", "TrafficOut"}
|
||||
hit := func(s string) string {
|
||||
for _, b := range banned {
|
||||
if strings.Contains(s, b) {
|
||||
return b
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
var violations []string
|
||||
ast.Inspect(f, func(n ast.Node) bool {
|
||||
switch v := n.(type) {
|
||||
case *ast.BasicLit:
|
||||
// SQL 列名 / gorm.Expr 里的字符串
|
||||
if v.Kind == token.STRING {
|
||||
if b := hit(v.Value); b != "" {
|
||||
violations = append(violations,
|
||||
fset.Position(v.Pos()).String()+" 出现字符串 "+b)
|
||||
}
|
||||
}
|
||||
case *ast.Ident:
|
||||
// 结构体字段 / 变量名
|
||||
if b := hit(v.Name); b != "" {
|
||||
violations = append(violations,
|
||||
fset.Position(v.Pos()).String()+" 出现标识符 "+b)
|
||||
}
|
||||
}
|
||||
return true
|
||||
})
|
||||
|
||||
if len(violations) > 0 {
|
||||
t.Fatalf("%s 又开始写流量列了,会和引擎侧计量双计:\n %s\n"+
|
||||
"流量只能由 waftask.FlushTrafficStats 写入,日志聚合只负责 count / time_spent。",
|
||||
file, strings.Join(violations, "\n "))
|
||||
}
|
||||
}
|
||||
|
||||
// 反向护栏:流量落库这条路径必须真的存在(防止被整体删掉后统计悄悄归零)
|
||||
func TestTrafficFlushPathExists(t *testing.T) {
|
||||
fset := token.NewFileSet()
|
||||
f, err := parser.ParseFile(fset, "task_traffic_stats.go", nil, 0)
|
||||
if err != nil {
|
||||
t.Fatalf("解析 task_traffic_stats.go 失败: %v", err)
|
||||
}
|
||||
|
||||
want := map[string]bool{
|
||||
"TaskTrafficFlush": false, // 定时任务入口
|
||||
"FlushTrafficStats": false, // 停机补刀也用它
|
||||
"writeTrafficStats": false,
|
||||
"planTrafficUpserts": false,
|
||||
}
|
||||
for _, decl := range f.Decls {
|
||||
if fn, ok := decl.(*ast.FuncDecl); ok {
|
||||
if _, exists := want[fn.Name.Name]; exists {
|
||||
want[fn.Name.Name] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
for name, found := range want {
|
||||
if !found {
|
||||
t.Fatalf("流量落库链路缺少函数 %s —— 统计会静默归零", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package waftask
|
||||
|
||||
import (
|
||||
"SamWaf/common/uuid"
|
||||
"SamWaf/common/zlog"
|
||||
"SamWaf/customtype"
|
||||
"SamWaf/global"
|
||||
"SamWaf/model"
|
||||
"SamWaf/model/baseorm"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// 站点流量落库:把引擎侧累加的真实进出字节按天/小时增量写进统计库。
|
||||
//
|
||||
// 与 CollectStatsFromLogs 的分工(issue #930 后):
|
||||
// - 本任务:只写 traffic_in / traffic_out,数据来自引擎字节计量,与日志无关;
|
||||
// - 日志聚合:只写 total/attack/normal_count 与 total_time_spent。
|
||||
//
|
||||
// 两边都用「先 Update 累加、影响 0 行再 Create」的写法,谁先建行都不冲突。
|
||||
|
||||
// trafficDayAgg 合并后的天级增量
|
||||
type trafficDayAgg struct {
|
||||
HostCode string
|
||||
Host string
|
||||
Day int
|
||||
In int64
|
||||
Out int64
|
||||
}
|
||||
|
||||
// trafficHourAgg 合并后的小时级增量
|
||||
type trafficHourAgg struct {
|
||||
HostCode string
|
||||
Host string
|
||||
HourTime int64
|
||||
In int64
|
||||
Out int64
|
||||
}
|
||||
|
||||
// planTrafficUpserts 把 Drain 出来的桶合并成天级/小时级两组增量。
|
||||
// 纯函数:同一天的多个整点桶会合并成一条天级增量,全零桶直接丢弃。
|
||||
func planTrafficUpserts(list []global.TrafficSnapshot) ([]trafficDayAgg, []trafficHourAgg) {
|
||||
type dayKey struct {
|
||||
HostCode string
|
||||
Day int
|
||||
}
|
||||
type hourKey struct {
|
||||
HostCode string
|
||||
HourTime int64
|
||||
}
|
||||
dayMap := make(map[dayKey]*trafficDayAgg)
|
||||
hourMap := make(map[hourKey]*trafficHourAgg)
|
||||
var days []*trafficDayAgg
|
||||
var hours []*trafficHourAgg
|
||||
|
||||
for _, s := range list {
|
||||
if s.HostCode == "" || (s.In <= 0 && s.Out <= 0) {
|
||||
continue
|
||||
}
|
||||
dk := dayKey{HostCode: s.HostCode, Day: s.Day}
|
||||
d := dayMap[dk]
|
||||
if d == nil {
|
||||
d = &trafficDayAgg{HostCode: s.HostCode, Host: s.Host, Day: s.Day}
|
||||
dayMap[dk] = d
|
||||
days = append(days, d)
|
||||
}
|
||||
d.In += s.In
|
||||
d.Out += s.Out
|
||||
|
||||
hk := hourKey{HostCode: s.HostCode, HourTime: s.HourTime}
|
||||
h := hourMap[hk]
|
||||
if h == nil {
|
||||
h = &trafficHourAgg{HostCode: s.HostCode, Host: s.Host, HourTime: s.HourTime}
|
||||
hourMap[hk] = h
|
||||
hours = append(hours, h)
|
||||
}
|
||||
h.In += s.In
|
||||
h.Out += s.Out
|
||||
}
|
||||
|
||||
dayList := make([]trafficDayAgg, 0, len(days))
|
||||
for _, d := range days {
|
||||
dayList = append(dayList, *d)
|
||||
}
|
||||
hourList := make([]trafficHourAgg, 0, len(hours))
|
||||
for _, h := range hours {
|
||||
hourList = append(hourList, *h)
|
||||
}
|
||||
return dayList, hourList
|
||||
}
|
||||
|
||||
// TaskTrafficFlush 定时任务入口(默认 30s 一次)
|
||||
func TaskTrafficFlush() {
|
||||
FlushTrafficStats()
|
||||
}
|
||||
|
||||
// FlushTrafficStats 把内存里累计的流量落库。
|
||||
// 库没就绪或正在切库时**不取走**内存里的增量,等下个周期,避免把字节丢在切库窗口里。
|
||||
func FlushTrafficStats() {
|
||||
if global.GWAF_LOCAL_STATS_DB == nil {
|
||||
return
|
||||
}
|
||||
if global.GDATA_CURRENT_CHANGE {
|
||||
zlog.Debug("流量统计落库", "正在切换数据库,本轮跳过")
|
||||
return
|
||||
}
|
||||
|
||||
list := global.DrainTraffic()
|
||||
if len(list) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
if err := writeTrafficStats(global.GWAF_LOCAL_STATS_DB, list); err != nil {
|
||||
// 整笔事务回滚了,把增量放回累加器下轮重试,绝不静默丢数
|
||||
global.RestoreTraffic(list)
|
||||
zlog.Error("流量统计落库失败,已退回内存等待重试", "错误", err.Error(), "桶数", len(list))
|
||||
return
|
||||
}
|
||||
zlog.Debug("流量统计落库完成", "桶数", len(list))
|
||||
}
|
||||
|
||||
// writeTrafficStats 单事务写入:要么全成,要么全退(避免部分成功后重试造成重复计数)
|
||||
func writeTrafficStats(db *gorm.DB, list []global.TrafficSnapshot) error {
|
||||
days, hours := planTrafficUpserts(list)
|
||||
if len(days) == 0 && len(hours) == 0 {
|
||||
return nil
|
||||
}
|
||||
now := customtype.JsonTime(time.Now())
|
||||
|
||||
return db.Transaction(func(tx *gorm.DB) error {
|
||||
for _, d := range days {
|
||||
res := tx.Model(&model.StatsSiteDay{}).
|
||||
Where("tenant_id = ? and user_code = ? and host_code = ? and day = ?",
|
||||
global.GWAF_TENANT_ID, global.GWAF_USER_CODE, d.HostCode, d.Day).
|
||||
Updates(map[string]interface{}{
|
||||
"traffic_in": gorm.Expr("traffic_in + ?", d.In),
|
||||
"traffic_out": gorm.Expr("traffic_out + ?", d.Out),
|
||||
"update_time": now,
|
||||
})
|
||||
if res.Error != nil {
|
||||
return res.Error
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
if err := tx.Create(&model.StatsSiteDay{
|
||||
BaseOrm: baseorm.BaseOrm{
|
||||
Id: uuid.GenUUID(),
|
||||
USER_CODE: global.GWAF_USER_CODE,
|
||||
Tenant_ID: global.GWAF_TENANT_ID,
|
||||
CREATE_TIME: now,
|
||||
UPDATE_TIME: now,
|
||||
},
|
||||
HostCode: d.HostCode,
|
||||
Day: d.Day,
|
||||
Host: d.Host,
|
||||
TrafficIn: d.In,
|
||||
TrafficOut: d.Out,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for _, h := range hours {
|
||||
res := tx.Model(&model.StatsSiteHour{}).
|
||||
Where("tenant_id = ? and user_code = ? and host_code = ? and hour_time = ?",
|
||||
global.GWAF_TENANT_ID, global.GWAF_USER_CODE, h.HostCode, h.HourTime).
|
||||
Updates(map[string]interface{}{
|
||||
"traffic_in": gorm.Expr("traffic_in + ?", h.In),
|
||||
"traffic_out": gorm.Expr("traffic_out + ?", h.Out),
|
||||
"update_time": now,
|
||||
})
|
||||
if res.Error != nil {
|
||||
return res.Error
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
if err := tx.Create(&model.StatsSiteHour{
|
||||
BaseOrm: baseorm.BaseOrm{
|
||||
Id: uuid.GenUUID(),
|
||||
USER_CODE: global.GWAF_USER_CODE,
|
||||
Tenant_ID: global.GWAF_TENANT_ID,
|
||||
CREATE_TIME: now,
|
||||
UPDATE_TIME: now,
|
||||
},
|
||||
HostCode: h.HostCode,
|
||||
HourTime: h.HourTime,
|
||||
Host: h.Host,
|
||||
TrafficIn: h.In,
|
||||
TrafficOut: h.Out,
|
||||
}).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,210 @@
|
||||
package waftask
|
||||
|
||||
import (
|
||||
"SamWaf/common/uuid"
|
||||
"SamWaf/global"
|
||||
"SamWaf/model"
|
||||
"SamWaf/model/baseorm"
|
||||
"net/url"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
sqlitedriver "github.com/samwafgo/sqlitedriver"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
)
|
||||
|
||||
// 真实库落盘验证:UPSERT 的「先累加、影响 0 行再新建」语义必须成立,
|
||||
// 否则要么第一笔流量丢失(没建行),要么每轮都新建行(重复计数)。
|
||||
// 无 CGO / 驱动不可用的环境自动跳过,不阻塞普通 go test。
|
||||
func openTrafficTestDB(t *testing.T) *gorm.DB {
|
||||
t.Helper()
|
||||
dsn := filepath.Join(t.TempDir(), "stats.db") + "?_db_key=" + url.QueryEscape("ktest")
|
||||
db, err := gorm.Open(sqlitedriver.Open(dsn), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Skipf("打不开 sqlite(缺 CGO?),跳过真实库用例: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.StatsSiteDay{}, &model.StatsSiteHour{}); err != nil {
|
||||
t.Skipf("建表失败,跳过真实库用例: %v", err)
|
||||
}
|
||||
// Windows 上不关连接,t.TempDir 清理会因文件被占用而报错
|
||||
t.Cleanup(func() {
|
||||
if sqlDB, e := db.DB(); e == nil {
|
||||
sqlDB.Close()
|
||||
}
|
||||
})
|
||||
return db
|
||||
}
|
||||
|
||||
func TestWriteTrafficStats_CreateThenAccumulate(t *testing.T) {
|
||||
db := openTrafficTestDB(t)
|
||||
|
||||
const day = 20260818
|
||||
const hour int64 = 1755500400
|
||||
snap := func(in, out int64) []global.TrafficSnapshot {
|
||||
return []global.TrafficSnapshot{{
|
||||
TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: day, HourTime: hour},
|
||||
In: in, Out: out,
|
||||
}}
|
||||
}
|
||||
|
||||
// 第一轮:表里没有行 → 应新建
|
||||
if err := writeTrafficStats(db, snap(100, 1000)); err != nil {
|
||||
t.Fatalf("首轮落库失败: %v", err)
|
||||
}
|
||||
var days []model.StatsSiteDay
|
||||
if err := db.Find(&days).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(days) != 1 {
|
||||
t.Fatalf("首轮应新建 1 行天记录,实际 %d 行", len(days))
|
||||
}
|
||||
if days[0].TrafficIn != 100 || days[0].TrafficOut != 1000 {
|
||||
t.Fatalf("首轮天流量 = %d/%d,期望 100/1000", days[0].TrafficIn, days[0].TrafficOut)
|
||||
}
|
||||
if days[0].TotalCount != 0 {
|
||||
t.Fatalf("流量落库不该动 PV 列,实际 total_count = %d", days[0].TotalCount)
|
||||
}
|
||||
|
||||
// 第二轮:同一天同整点 → 必须在原行上累加,而不是再插一行
|
||||
if err := writeTrafficStats(db, snap(50, 500)); err != nil {
|
||||
t.Fatalf("二轮落库失败: %v", err)
|
||||
}
|
||||
days = nil
|
||||
if err := db.Find(&days).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(days) != 1 {
|
||||
t.Fatalf("二轮不该新增行,实际 %d 行(重复计数)", len(days))
|
||||
}
|
||||
if days[0].TrafficIn != 150 || days[0].TrafficOut != 1500 {
|
||||
t.Fatalf("累加后天流量 = %d/%d,期望 150/1500", days[0].TrafficIn, days[0].TrafficOut)
|
||||
}
|
||||
|
||||
var hours []model.StatsSiteHour
|
||||
if err := db.Find(&hours).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(hours) != 1 {
|
||||
t.Fatalf("小时行应只有 1 行,实际 %d 行", len(hours))
|
||||
}
|
||||
if hours[0].TrafficIn != 150 || hours[0].TrafficOut != 1500 {
|
||||
t.Fatalf("累加后小时流量 = %d/%d,期望 150/1500", hours[0].TrafficIn, hours[0].TrafficOut)
|
||||
}
|
||||
}
|
||||
|
||||
// 已有日志聚合建好的行(有 PV 没流量)时,流量落库只能补流量列,不能覆盖 PV
|
||||
func TestWriteTrafficStats_DoesNotClobberExistingCounts(t *testing.T) {
|
||||
db := openTrafficTestDB(t)
|
||||
|
||||
const day = 20260818
|
||||
const hour int64 = 1755500400
|
||||
// 模拟 CollectStatsFromLogs 先建好行。
|
||||
// 租户/用户两列必须与 global 一致:流量 UPSERT 的 WHERE 带了这两列,
|
||||
// 对不上就会退化成"再插一行",天级数据被拆成两行。
|
||||
if err := db.Create(&model.StatsSiteDay{
|
||||
BaseOrm: baseorm.BaseOrm{
|
||||
Id: uuid.GenUUID(),
|
||||
USER_CODE: global.GWAF_USER_CODE,
|
||||
Tenant_ID: global.GWAF_TENANT_ID,
|
||||
},
|
||||
HostCode: "h1", Host: "a.com", Day: day,
|
||||
TotalCount: 218, AttackCount: 83, NormalCount: 135,
|
||||
}).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err := writeTrafficStats(db, []global.TrafficSnapshot{{
|
||||
TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: day, HourTime: hour},
|
||||
In: 7, Out: 9,
|
||||
}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var got model.StatsSiteDay
|
||||
if err := db.Where("host_code = ? and day = ?", "h1", day).First(&got).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got.TotalCount != 218 || got.AttackCount != 83 || got.NormalCount != 135 {
|
||||
t.Fatalf("流量落库把日志聚合的计数覆盖了: %+v", got)
|
||||
}
|
||||
if got.TrafficIn != 7 || got.TrafficOut != 9 {
|
||||
t.Fatalf("流量列没写上: %d/%d", got.TrafficIn, got.TrafficOut)
|
||||
}
|
||||
}
|
||||
|
||||
// 多站点多天一次落库:各行独立,不串账
|
||||
func TestWriteTrafficStats_MultiHostMultiDay(t *testing.T) {
|
||||
db := openTrafficTestDB(t)
|
||||
|
||||
list := []global.TrafficSnapshot{
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260818, HourTime: 1000}, In: 1, Out: 2},
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260818, HourTime: 4600}, In: 3, Out: 4},
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260819, HourTime: 90000}, In: 5, Out: 6},
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h2", Host: "b.com", Day: 20260818, HourTime: 1000}, In: 7, Out: 8},
|
||||
}
|
||||
if err := writeTrafficStats(db, list); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var days []model.StatsSiteDay
|
||||
if err := db.Order("host_code, day").Find(&days).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(days) != 3 {
|
||||
t.Fatalf("天级应为 3 行,实际 %d", len(days))
|
||||
}
|
||||
// h1 0818 两个整点合并
|
||||
if days[0].HostCode != "h1" || days[0].Day != 20260818 || days[0].TrafficIn != 4 || days[0].TrafficOut != 6 {
|
||||
t.Fatalf("h1 0818 天级合并错误: %+v", days[0])
|
||||
}
|
||||
if days[1].Day != 20260819 || days[1].TrafficIn != 5 {
|
||||
t.Fatalf("h1 0819 天级错误: %+v", days[1])
|
||||
}
|
||||
if days[2].HostCode != "h2" || days[2].TrafficIn != 7 {
|
||||
t.Fatalf("h2 天级错误: %+v", days[2])
|
||||
}
|
||||
|
||||
var hours []model.StatsSiteHour
|
||||
if err := db.Find(&hours).Error; err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(hours) != 4 {
|
||||
t.Fatalf("小时级应为 4 行(各整点分开),实际 %d", len(hours))
|
||||
}
|
||||
}
|
||||
|
||||
// 事务语义:中途失败必须整笔回滚,避免「一半写进去了」下轮重试造成重复计数
|
||||
func TestWriteTrafficStats_RollsBackOnError(t *testing.T) {
|
||||
db := openTrafficTestDB(t)
|
||||
|
||||
// 先写一笔正常数据
|
||||
if err := writeTrafficStats(db, []global.TrafficSnapshot{{
|
||||
TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260818, HourTime: 1000},
|
||||
In: 10, Out: 10,
|
||||
}}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 把小时表删掉,制造「天表能写、小时表必失败」的场景
|
||||
if err := db.Migrator().DropTable(&model.StatsSiteHour{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
err := writeTrafficStats(db, []global.TrafficSnapshot{{
|
||||
TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260818, HourTime: 1000},
|
||||
In: 999, Out: 999,
|
||||
}})
|
||||
if err == nil {
|
||||
t.Fatal("小时表不存在时应返回错误")
|
||||
}
|
||||
|
||||
var got model.StatsSiteDay
|
||||
if e := db.Where("host_code = ?", "h1").First(&got).Error; e != nil {
|
||||
t.Fatal(e)
|
||||
}
|
||||
if got.TrafficIn != 10 {
|
||||
t.Fatalf("事务没回滚:天表被写成 %d,期望仍是 10(否则退回内存重试会重复计数)", got.TrafficIn)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,105 @@
|
||||
package waftask
|
||||
|
||||
import (
|
||||
"SamWaf/global"
|
||||
"sort"
|
||||
"testing"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// planTrafficUpserts 是落库前的合并逻辑:同一天的多个整点桶要合成一条天级增量,
|
||||
// 小时级则必须保持分开——写错了就会出现「天总量对、小时曲线错」这种最难查的问题。
|
||||
func TestPlanTrafficUpserts_MergesDayKeepsHours(t *testing.T) {
|
||||
list := []global.TrafficSnapshot{
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260818, HourTime: 1000}, In: 10, Out: 100},
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260818, HourTime: 4600}, In: 5, Out: 50},
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260819, HourTime: 90000}, In: 1, Out: 2},
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h2", Host: "b.com", Day: 20260818, HourTime: 1000}, In: 7, Out: 8},
|
||||
}
|
||||
|
||||
days, hours := planTrafficUpserts(list)
|
||||
|
||||
if len(days) != 3 {
|
||||
t.Fatalf("天级增量应为 3 条(h1两天 + h2一天),实际 %d: %+v", len(days), days)
|
||||
}
|
||||
if len(hours) != 4 {
|
||||
t.Fatalf("小时级增量应为 4 条(各整点分开),实际 %d: %+v", len(hours), hours)
|
||||
}
|
||||
|
||||
sort.Slice(days, func(i, j int) bool {
|
||||
if days[i].HostCode != days[j].HostCode {
|
||||
return days[i].HostCode < days[j].HostCode
|
||||
}
|
||||
return days[i].Day < days[j].Day
|
||||
})
|
||||
// h1 的 0818:两个整点必须合并成 15/150
|
||||
if days[0].HostCode != "h1" || days[0].Day != 20260818 || days[0].In != 15 || days[0].Out != 150 {
|
||||
t.Fatalf("同一天的多个整点没合并对: %+v", days[0])
|
||||
}
|
||||
if days[0].Host != "a.com" {
|
||||
t.Fatalf("域名没带上: %+v", days[0])
|
||||
}
|
||||
if days[1].Day != 20260819 || days[1].In != 1 || days[1].Out != 2 {
|
||||
t.Fatalf("跨天的增量被并到一起了: %+v", days[1])
|
||||
}
|
||||
if days[2].HostCode != "h2" || days[2].In != 7 || days[2].Out != 8 {
|
||||
t.Fatalf("不同站点串账: %+v", days[2])
|
||||
}
|
||||
}
|
||||
|
||||
// 空桶/无 host_code 的记录不该产生任何写库动作
|
||||
func TestPlanTrafficUpserts_SkipsEmpty(t *testing.T) {
|
||||
list := []global.TrafficSnapshot{
|
||||
{TrafficKey: global.TrafficKey{HostCode: "", Host: "a.com", Day: 20260818, HourTime: 1000}, In: 10, Out: 10},
|
||||
{TrafficKey: global.TrafficKey{HostCode: "h1", Host: "a.com", Day: 20260818, HourTime: 1000}, In: 0, Out: 0},
|
||||
}
|
||||
days, hours := planTrafficUpserts(list)
|
||||
if len(days) != 0 || len(hours) != 0 {
|
||||
t.Fatalf("空桶不该落库,实际 days=%d hours=%d", len(days), len(hours))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPlanTrafficUpserts_Empty(t *testing.T) {
|
||||
days, hours := planTrafficUpserts(nil)
|
||||
if len(days) != 0 || len(hours) != 0 {
|
||||
t.Fatalf("空输入应返回空计划")
|
||||
}
|
||||
}
|
||||
|
||||
// 库没就绪时绝不能把内存里的增量取走丢掉——必须留着等下个周期
|
||||
func TestFlushTrafficStats_KeepsDataWhenDBNotReady(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
oldDB := global.GWAF_LOCAL_STATS_DB
|
||||
global.GWAF_LOCAL_STATS_DB = nil
|
||||
defer func() { global.GWAF_LOCAL_STATS_DB = oldDB }()
|
||||
|
||||
global.AddTraffic("h1", "a.com", 20260818, 1000, 123, 456)
|
||||
FlushTrafficStats()
|
||||
|
||||
if n := global.PendingTrafficBuckets(); n != 1 {
|
||||
t.Fatalf("库未就绪时应保留增量,实际待落库桶数 = %d(数据被丢了)", n)
|
||||
}
|
||||
global.DrainTraffic()
|
||||
}
|
||||
|
||||
// 切库窗口同理:跳过本轮,不取走
|
||||
func TestFlushTrafficStats_KeepsDataWhileSwitchingDB(t *testing.T) {
|
||||
global.DrainTraffic()
|
||||
oldDB := global.GWAF_LOCAL_STATS_DB
|
||||
// 必须给个非 nil 的库句柄,否则会在「库未就绪」那道判断就返回,测不到切库分支
|
||||
global.GWAF_LOCAL_STATS_DB = &gorm.DB{}
|
||||
global.GDATA_CURRENT_CHANGE = true
|
||||
defer func() {
|
||||
global.GDATA_CURRENT_CHANGE = false
|
||||
global.GWAF_LOCAL_STATS_DB = oldDB
|
||||
}()
|
||||
|
||||
global.AddTraffic("h1", "a.com", 20260818, 1000, 1, 1)
|
||||
FlushTrafficStats()
|
||||
|
||||
if n := global.PendingTrafficBuckets(); n != 1 {
|
||||
t.Fatalf("切库时应保留增量,实际待落库桶数 = %d", n)
|
||||
}
|
||||
global.DrainTraffic()
|
||||
}
|
||||
Reference in New Issue
Block a user