First Commit
ci / go (push) Waiting to run
ci / go-db (agent) (push) Waiting to run
ci / go-db (config) (push) Waiting to run
ci / go-db (db) (push) Waiting to run
ci / go-db (evidence) (push) Waiting to run
ci / go-db (llmrec) (push) Waiting to run
ci / go-db (server) (push) Waiting to run
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
detections / detections (push) Canceled after 0s
ci / go (push) Waiting to run
ci / go-db (agent) (push) Waiting to run
ci / go-db (config) (push) Waiting to run
ci / go-db (db) (push) Waiting to run
ci / go-db (evidence) (push) Waiting to run
ci / go-db (llmrec) (push) Waiting to run
ci / go-db (server) (push) Waiting to run
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
detections / detections (push) Canceled after 0s
This commit is contained in:
@@ -0,0 +1,975 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Autumn-27/artex/db"
|
||||
"github.com/Autumn-27/artex/notify"
|
||||
)
|
||||
|
||||
// 本文件覆盖推送功能的端到端行为:漏洞落库 → 事件 → 分派 → 真发 HTTP。
|
||||
//
|
||||
// 一个安全上的注意点:这些用例**不调用全局的 Notifier.step()**,只对自己创建的
|
||||
// 渠道调用 stepRealtime/stepDigest。原因是 step() 会遍历库里所有启用渠道——
|
||||
// 在一个已配好真实钉钉/企微机器人的开发库上跑测试,全局 step 会把测试期间
|
||||
// 产生的漏洞真推到那些群里。逐渠道调用把影响面严格限制在测试自造的假接收端上。
|
||||
//
|
||||
// 清理:用例结束时删掉本用例产生的事件(级联删投递)与渠道,不给真实渠道留积压。
|
||||
//
|
||||
// 断言口径:stepRealtime/stepDigest 不返回值、内部记日志,因此这里断言的是
|
||||
// **可观测的外部行为**(假接收端收到了什么、投递行落到什么状态),而不是函数的
|
||||
// 返回值——这比对返回值打桩更接近真实调用路径。
|
||||
|
||||
// notifyFixture 是本文件用例的公共装置。
|
||||
type notifyFixture struct {
|
||||
s *Server
|
||||
pg *db.DB
|
||||
request func(method, path, body string) *httptest.ResponseRecorder
|
||||
n *Notifier
|
||||
// 自建的 task/exploration:用例把漏洞记到这里,与其它用例的数据隔离。
|
||||
taskID int64
|
||||
expID int64
|
||||
// cleanupMark 之后产生的事件在清理时一并删除。
|
||||
cleanupMark int64
|
||||
}
|
||||
|
||||
func newNotifyFixture(t *testing.T) *notifyFixture {
|
||||
t.Helper()
|
||||
// 本文件的所有假接收端都跑在 127.0.0.1 上,而投递默认拒绝环回地址
|
||||
// (防 SSRF 打到同机服务与云元数据)。测试显式打开这个开关;
|
||||
// 守卫「默认拒绝」的行为由 notify 包的 ssrf_test.go 覆盖。
|
||||
t.Setenv(notify.AllowLocalTargetsEnv, "1")
|
||||
// 백그라운드 투递 루프(3초 tick)를 끈다. 이 파일의 케이스들은 stepRealtime·
|
||||
// stepDigest 를 직접 호출해 한 번의 분배 결과(앞 K건 송달·나머지 보류)를 검증하는데,
|
||||
// 백그라운드 루프가 같은 채널을 동시에 처리하면 집계가 타이밍에 따라 흔들려(느린 CI
|
||||
// 에서 간헐 실패) 결정성이 깨진다. 이 변수는 trafficEvidenceServer 가 서버를 세우기
|
||||
// 전에 설정돼야 효과가 있다.
|
||||
t.Setenv(notifyBackgroundDisabledEnv, "1")
|
||||
s, _, request := trafficEvidenceServer(t)
|
||||
pg := s.m.pg
|
||||
|
||||
// 自建一个 task:共享装置 trafficEvidenceServer 造的 task 拿不到 exploration id,
|
||||
// 而记录漏洞必须提供它。
|
||||
task, err := s.m.CreateTask("通知推送测试", "推送行为验证", nil, 0, 0)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
taskID, err := strconv.ParseInt(task.ID, 10, 64)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { pg.Exec(`DELETE FROM tasks WHERE id=$1`, taskID) })
|
||||
|
||||
var mark int64
|
||||
if err := pg.QueryRow(`SELECT COALESCE(max(id),0) FROM notification_events`).Scan(&mark); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 让装置自成闭环:把建 fixture 之前就存在的事件一次性标记为已分派。
|
||||
//
|
||||
// 为什么必须做:FanOutPendingEvents 是**全局**的,会把库里所有未分派事件
|
||||
// 展开到所有匹配渠道上。而共享装置 trafficEvidenceServer 自己就会记一条漏洞
|
||||
// (正是它返回的那个初始 finding),其它用例也可能有残留。不隔离的话,
|
||||
// 这些杂散事件会被分派到本用例的渠道上,让「应有 N 条投递」这类断言
|
||||
// 时对时错——而且错法取决于用例执行顺序,比直接失败更难查。
|
||||
if _, err := pg.Exec(`UPDATE notification_events SET fanned_out = true WHERE id <= $1 AND NOT fanned_out`, mark); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
f := ¬ifyFixture{s: s, pg: pg, request: request, n: newNotifier(s), taskID: taskID, expID: task.ExpID, cleanupMark: mark}
|
||||
t.Cleanup(func() {
|
||||
if _, err := pg.Exec(`DELETE FROM notification_events WHERE id > $1`, f.cleanupMark); err != nil {
|
||||
t.Logf("清理通知事件失败: %v", err)
|
||||
}
|
||||
})
|
||||
// 总开关必须是开的(其它用例可能关过它)。
|
||||
if err := pg.SetBool(settingNotifyEnabled, true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return f
|
||||
}
|
||||
|
||||
// record 走真实的证据写入路径落一条漏洞,返回 finding id。
|
||||
// 这条路会在**同一事务**里登记推送事件——正是本功能的挂点。
|
||||
func (f *notifyFixture) record(t *testing.T, vulnclass, severity string) int64 {
|
||||
t.Helper()
|
||||
out, err := f.s.evidenceStore().Record(context.Background(), db.RecordFindingInput{
|
||||
TaskID: f.taskID,
|
||||
ExplorationID: f.expID,
|
||||
Worker: "test",
|
||||
VulnClass: vulnclass,
|
||||
Name: vulnclass,
|
||||
Severity: severity,
|
||||
Summary: vulnclass + " 的摘要",
|
||||
Evidence: "poc",
|
||||
}, nil)
|
||||
if err != nil {
|
||||
t.Fatalf("记录漏洞失败: %v", err)
|
||||
}
|
||||
return out.FindingID
|
||||
}
|
||||
|
||||
// channel 读回渠道配置(供逐渠道调用 stepX 使用)。
|
||||
func (f *notifyFixture) channel(t *testing.T, id int64) *db.NotificationChannel {
|
||||
t.Helper()
|
||||
ch, err := f.pg.NotificationChannelByID(context.Background(), id)
|
||||
if err != nil {
|
||||
t.Fatalf("读渠道失败: %v", err)
|
||||
}
|
||||
return ch
|
||||
}
|
||||
|
||||
// deliver 分派事件并只对指定渠道跑一轮投递。
|
||||
func (f *notifyFixture) deliver(t *testing.T, chID int64, baseURL string) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
if _, _, err := f.pg.FanOutPendingEvents(ctx, 500); err != nil {
|
||||
t.Fatalf("分派失败: %v", err)
|
||||
}
|
||||
f.n.stepRealtime(ctx, f.channel(t, chID), 50, baseURL)
|
||||
}
|
||||
|
||||
// createChannel 通过 HTTP 接口建渠道,顺带覆盖接口自身的校验路径。
|
||||
func (f *notifyFixture) createChannel(t *testing.T, payload map[string]any) int64 {
|
||||
t.Helper()
|
||||
raw, _ := json.Marshal(payload)
|
||||
r := f.request("POST", "/api/notify/channels", string(raw))
|
||||
if r.Code != 200 {
|
||||
t.Fatalf("建渠道失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
var res struct {
|
||||
ID int64 `json:"id"`
|
||||
}
|
||||
if err := json.Unmarshal(r.Body.Bytes(), &res); err != nil || res.ID == 0 {
|
||||
t.Fatalf("建渠道返回异常: %s (%v)", r.Body, err)
|
||||
}
|
||||
t.Cleanup(func() { f.pg.Exec(`DELETE FROM notification_channels WHERE id=$1`, res.ID) })
|
||||
return res.ID
|
||||
}
|
||||
|
||||
// fakeWebhook 是记录收到的请求体的假接收端。
|
||||
type fakeWebhook struct {
|
||||
*httptest.Server
|
||||
mu sync.Mutex
|
||||
bodies []map[string]any
|
||||
}
|
||||
|
||||
func newFakeWebhook(t *testing.T) *fakeWebhook {
|
||||
t.Helper()
|
||||
f := &fakeWebhook{}
|
||||
f.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
raw, _ := io.ReadAll(r.Body)
|
||||
var body map[string]any
|
||||
_ = json.Unmarshal(raw, &body)
|
||||
f.mu.Lock()
|
||||
f.bodies = append(f.bodies, body)
|
||||
f.mu.Unlock()
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = io.WriteString(w, `{"errcode":0,"errmsg":"ok"}`)
|
||||
}))
|
||||
t.Cleanup(f.Close)
|
||||
return f
|
||||
}
|
||||
|
||||
func (f *fakeWebhook) count() int {
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
return len(f.bodies)
|
||||
}
|
||||
|
||||
func (f *fakeWebhook) body(t *testing.T, i int) map[string]any {
|
||||
t.Helper()
|
||||
f.mu.Lock()
|
||||
defer f.mu.Unlock()
|
||||
if i >= len(f.bodies) {
|
||||
t.Fatalf("假接收端只收到 %d 条请求,取不到第 %d 条", len(f.bodies), i)
|
||||
}
|
||||
return f.bodies[i]
|
||||
}
|
||||
|
||||
func (f *fakeWebhook) last(t *testing.T) map[string]any {
|
||||
t.Helper()
|
||||
if f.count() == 0 {
|
||||
t.Fatal("假接收端没有收到任何请求")
|
||||
}
|
||||
return f.body(t, f.count()-1)
|
||||
}
|
||||
|
||||
// markdownText 从请求体里取出正文,兼容各家的字段名差异:
|
||||
// 钉钉 markdown 用 `text`、ActionCard 用 `text`、企业微信 markdown 用 `content`。
|
||||
func markdownText(t *testing.T, body map[string]any) string {
|
||||
t.Helper()
|
||||
for _, key := range []string{"markdown", "actionCard"} {
|
||||
section, ok := body[key].(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
for _, field := range []string{"text", "content"} {
|
||||
if s, ok := section[field].(string); ok && s != "" {
|
||||
return s
|
||||
}
|
||||
}
|
||||
}
|
||||
t.Fatalf("请求体里没有可识别的正文: %v", body)
|
||||
return ""
|
||||
}
|
||||
|
||||
// agePendingBatch 把该渠道的待发投递催老,用于测试汇总批次到期。
|
||||
func (f *notifyFixture) agePendingBatch(t *testing.T, chID int64) {
|
||||
t.Helper()
|
||||
if _, err := f.pg.Exec(`UPDATE notification_deliveries SET created_at = now() - interval '2 hours'
|
||||
WHERE channel_id=$1 AND state=$2`, chID, db.NotifyStatePending); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyEndToEndRealtimeDelivery(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "实时推送",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
})
|
||||
f.record(t, "SQL注入", "high")
|
||||
f.deliver(t, chID, "")
|
||||
|
||||
if hook.count() != 1 {
|
||||
t.Fatalf("应发出 1 条消息,实际 %d", hook.count())
|
||||
}
|
||||
text := markdownText(t, hook.last(t))
|
||||
// "SQL注入" 은 입력으로 넣은 취약점 제목이라 카드에 그대로 에코된다(사용자 데이터,
|
||||
// 번역 대상 아님). "높음"·"개요" 는 렌더 라벨이다 — 심각도는 notify.SeverityLabel("high")
|
||||
// 가 내는 "🟠 높음", 요약 머리글은 writeItem 이 붙이는 "**개요**:" 에 각각 들어 있다.
|
||||
for _, want := range []string{"SQL注入", "높음", "개요"} {
|
||||
if !strings.Contains(text, want) {
|
||||
t.Fatalf("消息正文缺少 %q:\n%s", want, text)
|
||||
}
|
||||
}
|
||||
// 投递应流转为 sent。
|
||||
var pending int
|
||||
if err := f.pg.QueryRow(`SELECT count(*) FROM notification_deliveries WHERE channel_id=$1 AND state <> $2`,
|
||||
chID, db.NotifyStateSent).Scan(&pending); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if pending != 0 {
|
||||
t.Fatalf("投递后仍有 %d 条未标记 sent", pending)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyChannelAPIMasksSecretsAndPreservesOnUpdate(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "掩码用例",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": "https://oapi.dingtalk.com/robot/send?access_token=abc123456", "secret": "SECabcdef123456"},
|
||||
})
|
||||
|
||||
r := f.request("GET", "/api/notify/channels", "")
|
||||
if r.Code != 200 {
|
||||
t.Fatalf("列渠道失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
if strings.Contains(r.Body.String(), "abc123456") || strings.Contains(r.Body.String(), "SECabcdef123456") {
|
||||
t.Fatalf("接口回显泄露了凭据: %s", r.Body)
|
||||
}
|
||||
var listed struct {
|
||||
Channels []struct {
|
||||
ID int64 `json:"id"`
|
||||
Config map[string]any `json:"config"`
|
||||
SecretKeys []string `json:"secret_keys"`
|
||||
} `json:"channels"`
|
||||
}
|
||||
if err := json.Unmarshal(r.Body.Bytes(), &listed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var mine *struct {
|
||||
ID int64 `json:"id"`
|
||||
Config map[string]any `json:"config"`
|
||||
SecretKeys []string `json:"secret_keys"`
|
||||
}
|
||||
for i := range listed.Channels {
|
||||
if listed.Channels[i].ID == chID {
|
||||
mine = &listed.Channels[i]
|
||||
}
|
||||
}
|
||||
if mine == nil {
|
||||
t.Fatal("新建的渠道未出现在列表里")
|
||||
}
|
||||
if !notify.IsMasked(fmt.Sprint(mine.Config["webhook"])) || !notify.IsMasked(fmt.Sprint(mine.Config["secret"])) {
|
||||
t.Fatalf("凭据字段应为掩码值: %v", mine.Config)
|
||||
}
|
||||
if len(mine.SecretKeys) == 0 {
|
||||
t.Fatal("接口应告知前端哪些字段是凭据")
|
||||
}
|
||||
|
||||
// PATCH 只改名 + 回传掩码凭据:真凭据必须原样保留。
|
||||
body, _ := json.Marshal(map[string]any{
|
||||
"name": "改名后",
|
||||
"config": map[string]any{"webhook": fmt.Sprint(mine.Config["webhook"]), "secret": fmt.Sprint(mine.Config["secret"])},
|
||||
})
|
||||
if r := f.request("PATCH", fmt.Sprintf("/api/notify/channels/%d", chID), string(body)); r.Code != 200 {
|
||||
t.Fatalf("更新失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
cfg := f.channelConfig(t, chID)
|
||||
if cfg["webhook"] != "https://oapi.dingtalk.com/robot/send?access_token=abc123456" {
|
||||
t.Fatalf("掩码回传把真凭据覆盖了: %v", cfg["webhook"])
|
||||
}
|
||||
if cfg["secret"] != "SECabcdef123456" {
|
||||
t.Fatalf("掩码回传把 secret 覆盖了: %v", cfg["secret"])
|
||||
}
|
||||
if f.channel(t, chID).Name != "改名后" {
|
||||
t.Fatal("名字未更新")
|
||||
}
|
||||
|
||||
// 显式清空 secret 应生效(区别于「回传掩码=保持不变」)。
|
||||
body, _ = json.Marshal(map[string]any{"config": map[string]any{"secret": ""}})
|
||||
if r := f.request("PATCH", fmt.Sprintf("/api/notify/channels/%d", chID), string(body)); r.Code != 200 {
|
||||
t.Fatalf("清空 secret 失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
if _, still := f.channelConfig(t, chID)["secret"]; still {
|
||||
t.Fatal("空串应清空 secret")
|
||||
}
|
||||
}
|
||||
|
||||
func (f *notifyFixture) channelConfig(t *testing.T, id int64) map[string]any {
|
||||
t.Helper()
|
||||
var cfg map[string]any
|
||||
if err := json.Unmarshal(f.channel(t, id).Config, &cfg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return cfg
|
||||
}
|
||||
|
||||
func TestNotifyChannelAPICreateValidation(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
cases := []struct {
|
||||
name string
|
||||
payload map[string]any
|
||||
wantSub string
|
||||
}{
|
||||
// wantSub 는 server/notify_api.go 가 내보내는 한국어 검증 오류와 일치시킨다
|
||||
// (notifyErrKindInvalidFmt·notifyErrNameMissing·notifyErrModeInvalid 의 안정 문구).
|
||||
// webhook 두 건은 notify/dingtalk.go 의 Validate 가 내는 오류인데, 그 webhook 주소
|
||||
// 검증 문구는 이미 한국어로 현지화됐으므로(notify/dingtalk_feishu_wecom_localized_test.go
|
||||
// 가 "Webhook 주소가 없습니다"·"Webhook 주소가 올바르지 않습니다" 로 검증한다) 기대
|
||||
// 문자열도 한국어 안정 문구로 맞춘다.
|
||||
{"类型非法", map[string]any{"name": "x", "kind": "nope", "config": map[string]any{}}, "채널 유형이 올바르지 않습니다"},
|
||||
{"缺名称", map[string]any{"kind": notify.KindDingTalk, "config": map[string]any{"webhook": "https://e.com/h"}}, "채널 이름을 입력하세요"},
|
||||
{"缺 webhook", map[string]any{"name": "x", "kind": notify.KindDingTalk, "config": map[string]any{}}, "Webhook 주소가 없습니다"},
|
||||
{"webhook 协议非法", map[string]any{"name": "x", "kind": notify.KindDingTalk, "config": map[string]any{"webhook": "file:///etc/passwd"}}, "Webhook 주소가 올바르지 않습니다"},
|
||||
{"模式非法", map[string]any{"name": "x", "kind": notify.KindDingTalk, "mode": "sometimes", "config": map[string]any{"webhook": "https://e.com/h"}}, "발송 모드가 올바르지 않습니다"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
raw, _ := json.Marshal(tc.payload)
|
||||
r := f.request("POST", "/api/notify/channels", string(raw))
|
||||
if r.Code != 400 {
|
||||
t.Fatalf("应返回 400,得到 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
if !strings.Contains(r.Body.String(), tc.wantSub) {
|
||||
t.Fatalf("错误信息应提到 %q,得到 %s", tc.wantSub, r.Body)
|
||||
}
|
||||
})
|
||||
}
|
||||
if r := f.request("DELETE", "/api/notify/channels/99999999", ""); r.Code != 404 {
|
||||
t.Fatalf("删除不存在的渠道应 404,得到 %d", r.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyFilterBlocksBelowThreshold(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "仅严重",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
"filter": map[string]any{"min_severity": "critical"},
|
||||
})
|
||||
f.record(t, "低危问题", "low")
|
||||
if _, _, err := f.pg.FanOutPendingEvents(context.Background(), 500); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
if err := f.pg.QueryRow(`SELECT count(*) FROM notification_deliveries WHERE channel_id=$1`, chID).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("低于阈值的漏洞不该产生投递,得到 %d 条", n)
|
||||
}
|
||||
f.n.stepRealtime(context.Background(), f.channel(t, chID), 50, "")
|
||||
if hook.count() != 0 {
|
||||
t.Fatal("被过滤的漏洞不应发出消息")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyDigestBatchesMultipleFindingsIntoOneMessage(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "汇总推送",
|
||||
"kind": notify.KindDingTalk,
|
||||
"mode": db.NotifyModeDigest,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
})
|
||||
for i := 0; i < 3; i++ {
|
||||
f.record(t, fmt.Sprintf("汇总漏洞%d", i+1), "high")
|
||||
}
|
||||
ctx := context.Background()
|
||||
if _, _, err := f.pg.FanOutPendingEvents(ctx, 500); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ch := f.channel(t, chID)
|
||||
|
||||
// 未到期:不发。
|
||||
f.n.stepDigest(ctx, ch, 50, "")
|
||||
if hook.count() != 0 {
|
||||
t.Fatal("汇总批次未到期就发了")
|
||||
}
|
||||
|
||||
// 催老批次后:三条合成一条消息。
|
||||
f.agePendingBatch(t, chID)
|
||||
f.n.stepDigest(ctx, ch, 50, "")
|
||||
if got := hook.count(); got != 1 {
|
||||
t.Fatalf("三条应汇总成一条消息,实际发了 %d 条", got)
|
||||
}
|
||||
text := markdownText(t, hook.last(t))
|
||||
// 시간창이 있는 다이제스트 머리말은 markdown.go 가 "**최근 N분간 신규 취약점 N건**"
|
||||
// 으로 렌더한다(digestInterval 기본 30분이라 WindowMinutes>0).
|
||||
if !strings.Contains(text, "최근") || !strings.Contains(text, "신규 취약점 3건") {
|
||||
t.Fatalf("汇总消息缺少条数/时间窗文案:\n%s", text)
|
||||
}
|
||||
for i := 1; i <= 3; i++ {
|
||||
if !strings.Contains(text, fmt.Sprintf("汇总漏洞%d", i)) {
|
||||
t.Fatalf("汇总消息缺少第 %d 条:\n%s", i, text)
|
||||
}
|
||||
}
|
||||
// 同一批次应共享 batch_id。
|
||||
var distinct, total int
|
||||
if err := f.pg.QueryRow(`SELECT count(DISTINCT batch_id), count(*) FROM notification_deliveries WHERE channel_id=$1`, chID).Scan(&distinct, &total); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if total != 3 || distinct != 1 {
|
||||
t.Fatalf("三条投递应共享一个 batch_id,得到 distinct=%d total=%d", distinct, total)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyDisabledChannelDoesNotSend(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "停用渠道",
|
||||
"kind": notify.KindDingTalk,
|
||||
"enabled": false,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
})
|
||||
f.record(t, "停用期间的漏洞", "critical")
|
||||
if _, _, err := f.pg.FanOutPendingEvents(context.Background(), 500); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
if err := f.pg.QueryRow(`SELECT count(*) FROM notification_deliveries WHERE channel_id=$1`, chID).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("停用渠道不该产生投递,得到 %d 条", n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyStatusChangeDelivery(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "状态变更订阅",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
"filter": map[string]any{"on_status_change": true},
|
||||
})
|
||||
finding := f.record(t, "状态变更用例", "high")
|
||||
r := f.request("PATCH", fmt.Sprintf("/api/exploration/findings/%d", finding), `{"status":"fixed"}`)
|
||||
if r.Code != 200 {
|
||||
t.Fatalf("改状态失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
f.deliver(t, chID, "")
|
||||
|
||||
// 应有两条:fixed 那一条是状态变更;finding_created 那条也可能在同一轮发出。
|
||||
// 状态变更的实际上更晚创建,但不依赖顺序,全量找。
|
||||
found := false
|
||||
for i := 0; i < hook.count(); i++ {
|
||||
text := markdownText(t, hook.body(t, i))
|
||||
// 상태 변경 카드는 markdown.go 가 "**상태 변경**: %s → %s" 로, 상태값은
|
||||
// notify.StatusLabel 이 렌더한다(fixed → "수정됨").
|
||||
if strings.Contains(text, "상태 변경") && strings.Contains(text, "수정됨") {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("没有收到含「状态变更 → 已修复」的消息(共 %d 条)", hook.count())
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyStatusChangeSuppressedByDefault(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "不订阅状态变更",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
})
|
||||
finding := f.record(t, "不订阅变更", "high")
|
||||
if r := f.request("PATCH", fmt.Sprintf("/api/exploration/findings/%d", finding), `{"status":"false_positive"}`); r.Code != 200 {
|
||||
t.Fatalf("改状态失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
if _, _, err := f.pg.FanOutPendingEvents(context.Background(), 500); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var n int
|
||||
if err := f.pg.QueryRow(`SELECT count(*) FROM notification_deliveries d
|
||||
JOIN notification_events e ON e.id = d.event_id
|
||||
WHERE d.channel_id=$1 AND e.kind=$2`, chID, notify.EventFindingStatusChanged).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("未订阅状态变更的渠道不该收到状态变更投递,得到 %d 条", n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyTestMessageEndpoint(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "测试发送",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
})
|
||||
if r := f.request("POST", fmt.Sprintf("/api/notify/channels/%d/test", chID), ""); r.Code != 200 {
|
||||
t.Fatalf("测试发送失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
if hook.count() != 1 {
|
||||
t.Fatalf("假接收端应收到 1 条测试消息,得到 %d", hook.count())
|
||||
}
|
||||
// 测试消息必须一眼能看出是测试,不能被误当成真实漏洞。
|
||||
// 테스트 메시지 제목은 notify_api.go 의 notifyTestName("테스트 메시지 · 채널 설정 정상").
|
||||
if text := markdownText(t, hook.last(t)); !strings.Contains(text, "테스트") {
|
||||
t.Fatalf("测试消息应标明是测试: %s", text)
|
||||
}
|
||||
// 配置坏掉时应把渠道的原始错误如实回给用户。
|
||||
badID := f.createChannel(t, map[string]any{
|
||||
"name": "坏地址",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": "http://127.0.0.1:1/hook"},
|
||||
})
|
||||
if r := f.request("POST", fmt.Sprintf("/api/notify/channels/%d/test", badID), ""); r.Code != 502 {
|
||||
t.Fatalf("投递失败应回 502,得到 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyDeliveriesHistoryAndRetry(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
// 指向必然失败的地址,制造 failed 投递。
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "失败重试",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": "http://127.0.0.1:1/hook"},
|
||||
})
|
||||
f.record(t, "会失败的推送", "high")
|
||||
ctx := context.Background()
|
||||
if _, _, err := f.pg.FanOutPendingEvents(ctx, 500); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ch := f.channel(t, chID)
|
||||
// 连投到耗尽重试预算。
|
||||
for i := 0; i < db.MaxNotifyAttempts; i++ {
|
||||
f.n.stepRealtime(ctx, ch, 50, "")
|
||||
if _, err := f.pg.Exec(`UPDATE notification_deliveries SET next_attempt_at = now() - interval '1 minute' WHERE channel_id=$1`, chID); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
var state string
|
||||
if err := f.pg.QueryRow(`SELECT state FROM notification_deliveries WHERE channel_id=$1`, chID).Scan(&state); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if state != db.NotifyStateFailed {
|
||||
t.Fatalf("重试耗尽后应为 failed,得到 %s", state)
|
||||
}
|
||||
|
||||
r := f.request("GET", fmt.Sprintf("/api/notify/deliveries?channel_id=%d&state=failed", chID), "")
|
||||
if r.Code != 200 {
|
||||
t.Fatalf("查历史失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
var hist struct {
|
||||
Deliveries []struct {
|
||||
ID int64 `json:"id"`
|
||||
State string `json:"state"`
|
||||
LastError string `json:"last_error"`
|
||||
Attempts int `json:"attempts"`
|
||||
Title string `json:"title"`
|
||||
} `json:"deliveries"`
|
||||
Total int `json:"total"`
|
||||
}
|
||||
if err := json.Unmarshal(r.Body.Bytes(), &hist); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hist.Total != 1 || len(hist.Deliveries) != 1 {
|
||||
t.Fatalf("应查到 1 条失败投递,得到 total=%d len=%d", hist.Total, len(hist.Deliveries))
|
||||
}
|
||||
if hist.Deliveries[0].LastError == "" {
|
||||
t.Fatal("历史里应带上失败原因,否则用户无法排查")
|
||||
}
|
||||
if hist.Deliveries[0].Attempts < db.MaxNotifyAttempts {
|
||||
t.Fatalf("尝试次数应被记录,得到 %d", hist.Deliveries[0].Attempts)
|
||||
}
|
||||
if hist.Deliveries[0].Title != "会失败的推送" {
|
||||
t.Fatalf("历史应带出漏洞标题,得到 %q", hist.Deliveries[0].Title)
|
||||
}
|
||||
|
||||
// 手动重发:应回到 pending 且清零计数。
|
||||
if r := f.request("POST", fmt.Sprintf("/api/notify/deliveries/%d/retry", hist.Deliveries[0].ID), ""); r.Code != 200 {
|
||||
t.Fatalf("重发失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
var attempts int
|
||||
if err := f.pg.QueryRow(`SELECT state, attempts FROM notification_deliveries WHERE id=$1`, hist.Deliveries[0].ID).Scan(&state, &attempts); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if state != db.NotifyStatePending || attempts != 0 {
|
||||
t.Fatalf("重发后应为 pending 且 attempts=0,得到 %s/%d", state, attempts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotifyMetaAndSettingsRoundTrip(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
r := f.request("GET", "/api/notify/meta", "")
|
||||
if r.Code != 200 {
|
||||
t.Fatalf("meta 失败: %s", r.Body)
|
||||
}
|
||||
var meta struct {
|
||||
Kinds []struct {
|
||||
Kind string `json:"kind"`
|
||||
SecretKeys []string `json:"secret_keys"`
|
||||
} `json:"kinds"`
|
||||
}
|
||||
if err := json.Unmarshal(r.Body.Bytes(), &meta); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(meta.Kinds) != len(notify.Kinds()) {
|
||||
t.Fatalf("meta 应列出全部 %d 个渠道,得到 %d", len(notify.Kinds()), len(meta.Kinds))
|
||||
}
|
||||
for _, k := range meta.Kinds {
|
||||
if len(k.SecretKeys) == 0 {
|
||||
t.Errorf("渠道 %s 未上报凭据字段", k.Kind)
|
||||
}
|
||||
}
|
||||
|
||||
// 三项全局设置往返。尾部斜杠应被规范化掉,否则回链会拼出 "//function/..."。
|
||||
if r := f.request("PUT", "/api/settings", `{"notify_public_base_url":"https://artex.example.com/","notify_digest_interval_min":15,"notify_enabled":true}`); r.Code != 200 {
|
||||
t.Fatalf("写设置失败 %d: %s", r.Code, r.Body)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
f.pg.Exec(`DELETE FROM settings WHERE key IN ($1,$2)`, settingNotifyPublicBaseURL, settingNotifyDigestMinutes)
|
||||
})
|
||||
payload := f.s.settingsPayload()
|
||||
if payload["notify_public_base_url"] != "https://artex.example.com" {
|
||||
t.Fatalf("回链地址未规范化: %v", payload["notify_public_base_url"])
|
||||
}
|
||||
if payload["notify_digest_interval_min"] != 15 {
|
||||
t.Fatalf("汇总周期未生效: %v", payload["notify_digest_interval_min"])
|
||||
}
|
||||
|
||||
// 非法值应被拒。
|
||||
for _, body := range []string{
|
||||
`{"notify_public_base_url":"ftp://x"}`,
|
||||
`{"notify_digest_interval_min":0}`,
|
||||
`{"notify_digest_interval_min":99999}`,
|
||||
} {
|
||||
if r := f.request("PUT", "/api/settings", body); r.Code != 400 {
|
||||
t.Errorf("%s 应返回 400,得到 %d", body, r.Code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyDeepLinkUsesPublicBaseURL 覆盖回链拼接:配了 public_base_url 时
|
||||
// 单条消息必须用带按钮的 ActionCard,且链接指向漏洞详情页。
|
||||
func TestNotifyDeepLinkUsesPublicBaseURL(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "回链",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
})
|
||||
finding := f.record(t, "带回链的漏洞", "high")
|
||||
f.deliver(t, chID, "https://artex.example.com")
|
||||
|
||||
body := hook.last(t)
|
||||
card, _ := body["actionCard"].(map[string]any)
|
||||
if card == nil {
|
||||
t.Fatalf("有回链时应用 ActionCard,得到 msgtype=%v", body["msgtype"])
|
||||
}
|
||||
want := fmt.Sprintf("https://artex.example.com/function/findings/detail?id=%d", finding)
|
||||
if card["singleURL"] != want {
|
||||
t.Fatalf("回链不对\n期望 %s\n得到 %v", want, card["singleURL"])
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyNoDeepLinkWithoutBaseURL 反向覆盖:没配外部地址时不该产生坏链接
|
||||
// (比如指向 localhost 或相对路径),应退回纯 markdown。
|
||||
func TestNotifyNoDeepLinkWithoutBaseURL(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "无回链",
|
||||
"kind": notify.KindDingTalk,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
})
|
||||
f.record(t, "无回链的漏洞", "high")
|
||||
f.deliver(t, chID, "")
|
||||
|
||||
body := hook.last(t)
|
||||
if body["msgtype"] != "markdown" {
|
||||
t.Fatalf("未配外部地址时应发 markdown,得到 %v", body["msgtype"])
|
||||
}
|
||||
// 상세 링크 머리글은 markdown.go 가 "[상세 보기](URL)" 로 렌더한다 — 외부 주소가
|
||||
// 없으면 이 링크가 아예 나오지 않아야 한다(옛 중국어 "查看详情" 를 검사하면 라벨이
|
||||
// 한국어로 바뀐 지금은 항상 통과해 회귀를 못 잡는다).
|
||||
if text := markdownText(t, body); strings.Contains(text, "상세 보기") {
|
||||
t.Fatalf("未配外部地址时不该出现详情链接:\n%s", text)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyDigestSegmentsAndDefersRemainder 是「静默丢失」修复的端到端证据。
|
||||
//
|
||||
// 汇总消息受渠道长度上限约束(企微 4096 字节),一批装不下时必须**按整条**切分:
|
||||
// 装进本条的那些标记已送达,其余回到队列等下一条。曾经的实现是把整批标记
|
||||
// 成功——被截掉的那些既不在消息里、也不在失败列表里,投递历史还显示成功,
|
||||
// 漏洞就这么没了。
|
||||
//
|
||||
// 断言四件事:① 只标记了实际装下的条数 ② 其余仍是待发 ③ 被推迟的条目
|
||||
// **没有消耗重试次数** ④ 再跑一轮能把剩下的发出去(不会卡死)。
|
||||
func TestNotifyDigestSegmentsAndDefersRemainder(t *testing.T) {
|
||||
f := newNotifyFixture(t)
|
||||
hook := newFakeWebhook(t)
|
||||
// 用企业微信:markdown 上限 4096 字节,是六个渠道里最紧的。
|
||||
chID := f.createChannel(t, map[string]any{
|
||||
"name": "分段汇总",
|
||||
"kind": notify.KindWeCom,
|
||||
"mode": db.NotifyModeDigest,
|
||||
"config": map[string]any{"webhook": hook.URL},
|
||||
})
|
||||
const total = 60
|
||||
// 标题取长一点,保证 60 条远超 4096 字节,必然分段。
|
||||
longName := strings.Repeat("超长漏洞名称", 6)
|
||||
for i := 0; i < total; i++ {
|
||||
f.record(t, longName+strconv.Itoa(i+1), "high")
|
||||
}
|
||||
ctx := context.Background()
|
||||
if _, _, err := f.pg.FanOutPendingEvents(ctx, 500); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
f.agePendingBatch(t, chID)
|
||||
ch := f.channel(t, chID)
|
||||
|
||||
f.n.stepDigest(ctx, ch, 50, "")
|
||||
if hook.count() != 1 {
|
||||
t.Fatalf("应只发出一条消息,得到 %d", hook.count())
|
||||
}
|
||||
|
||||
var sent, pending int
|
||||
if err := f.pg.QueryRow(`SELECT
|
||||
count(*) FILTER (WHERE state=$2),
|
||||
count(*) FILTER (WHERE state=$3)
|
||||
FROM notification_deliveries WHERE channel_id=$1`, chID, db.NotifyStateSent, db.NotifyStatePending).
|
||||
Scan(&sent, &pending); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if sent == 0 {
|
||||
t.Fatal("应有条目被标记为已送达")
|
||||
}
|
||||
if pending == 0 {
|
||||
t.Fatalf("一批 %d 条不可能全装进 4096 字节,应有剩余待发;sent=%d", total, sent)
|
||||
}
|
||||
if sent+pending != total {
|
||||
t.Fatalf("条目数对不上:sent=%d pending=%d total=%d(既没送达也没待发=丢失)", sent, pending, total)
|
||||
}
|
||||
// 消息正文必须如实告知还有多少条没包含在本条里。
|
||||
// 분절 안내는 markdown.go 가 "(이 메시지에는 앞 N건만 … 나머지 N건은 다음 메시지에서 …)" 로 렌더한다.
|
||||
if text := markdownText(t, hook.last(t)); !strings.Contains(text, "나머지") {
|
||||
t.Fatalf("消息应说明还有条目未包含在本条:\n%.400s", text)
|
||||
}
|
||||
|
||||
// 被推迟的条目不得消耗重试预算:领取时 attempts 已乐观 +1,推迟时要减回去。
|
||||
var maxAttempts int
|
||||
if err := f.pg.QueryRow(`SELECT COALESCE(max(attempts),0) FROM notification_deliveries
|
||||
WHERE channel_id=$1 AND state=$2`, chID, db.NotifyStatePending).Scan(&maxAttempts); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if maxAttempts > 0 {
|
||||
t.Fatalf("被推迟的条目不该消耗重试次数(否则几条之后就会被判失败),得到 attempts=%d", maxAttempts)
|
||||
}
|
||||
|
||||
// 反复跑直到收敛。断言的是**最终全部送达**且中途确实分了多轮——
|
||||
// 这比「第二轮发完」更强:它证明分段不会卡死、也不会把剩余条目丢掉。
|
||||
rounds := 0
|
||||
for {
|
||||
var undelivered int
|
||||
if err := f.pg.QueryRow(`SELECT count(*) FROM notification_deliveries
|
||||
WHERE channel_id=$1 AND state <> $2 AND state <> $3`, chID, db.NotifyStateSent, db.NotifyStateFailed).
|
||||
Scan(&undelivered); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if undelivered == 0 {
|
||||
break
|
||||
}
|
||||
rounds++
|
||||
if rounds > total+5 {
|
||||
t.Fatalf("分段投递不收敛:跑了 %d 轮仍有 %d 条悬而未决", rounds, undelivered)
|
||||
}
|
||||
before := hook.count()
|
||||
f.n.stepDigest(ctx, ch, 50, "")
|
||||
if hook.count() == before {
|
||||
t.Fatalf("第 %d 轮没有任何进展,剩余 %d 条会永久卡住", rounds, undelivered)
|
||||
}
|
||||
}
|
||||
if rounds < 2 {
|
||||
t.Fatalf("一条 4096 字节的消息装不下 %d 条长标题漏洞,应分多轮发出,实际只用了 %d 轮", total, rounds)
|
||||
}
|
||||
// 首轮之后的每一轮都应是**纯续发**,不存在被渠道拒绝的条目。
|
||||
var failed int
|
||||
if err := f.pg.QueryRow(`SELECT count(*) FROM notification_deliveries WHERE channel_id=$1 AND state=$2`,
|
||||
chID, db.NotifyStateFailed).Scan(&failed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if failed != 0 {
|
||||
t.Fatalf("假接收端始终返回成功,不该有失败条目,得到 %d", failed)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyBackoffTableMatchesAttemptBudget 是防漂移断言。
|
||||
//
|
||||
// 重试预算(db.MaxNotifyAttempts)与退避序列表(notifyBackoff)分居两个包:
|
||||
// 前者是状态机的策略、后者是引擎的执行节拍。若只改其中一个——比如把预算提到 5
|
||||
// 次而忘了加退避档位——代码不会报错,只会让第 4、5 次重试沿用最后一档间隔,
|
||||
// 表现为「重试节奏莫名变慢」,排查时很难联想到是这里。
|
||||
// 断言两者长度一致,让这种漂移在 CI 里就暴露。
|
||||
func TestNotifyBackoffTableMatchesAttemptBudget(t *testing.T) {
|
||||
if len(notifyBackoff) != db.MaxNotifyAttempts {
|
||||
t.Fatalf("退避档位数(%d)与最大尝试次数(%d)不一致——改一个必须同时改另一个",
|
||||
len(notifyBackoff), db.MaxNotifyAttempts)
|
||||
}
|
||||
// 退避间隔必须单调不减,否则重试会越试越急,反而加剧限流。
|
||||
for i := 1; i < len(notifyBackoff); i++ {
|
||||
if notifyBackoff[i] < notifyBackoff[i-1] {
|
||||
t.Fatalf("退避间隔必须单调不减:第 %d 档 %v < 第 %d 档 %v",
|
||||
i, notifyBackoff[i], i-1, notifyBackoff[i-1])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyRateLimitDoesNotConsumeRetryBudget 锁住「先取令牌再领取」的顺序。
|
||||
// 若反了(先领后弃),被限流挡下的投递已经计过一次 attempts,
|
||||
// 预算会被纯粹的等待耗光,最后落进 failed。
|
||||
func TestNotifyRateLimitDoesNotConsumeRetryBudget(t *testing.T) {
|
||||
// 只测令牌桶本身,不需要 Server(也不该为它造一个)。
|
||||
n := &Notifier{buckets: map[int64]*notifyBucket{}}
|
||||
now := time.Now()
|
||||
// 每分钟 1 条:满桶时最多 1 条。
|
||||
if got := n.takeTokens(1, 1, notifyMaxSendsPerChannelPerTick, now); got != 1 {
|
||||
t.Fatalf("满桶时每分钟 1 条应取 1 个令牌,得到 %d", got)
|
||||
}
|
||||
if got := n.takeTokens(1, 1, notifyMaxSendsPerChannelPerTick, now.Add(time.Millisecond)); got != 0 {
|
||||
t.Fatalf("令牌耗尽后应立即返回 0,得到 %d", got)
|
||||
}
|
||||
if got := n.takeTokens(1, 1, notifyMaxSendsPerChannelPerTick, now.Add(30*time.Second)); got != 0 {
|
||||
t.Fatalf("半程不应补满一个令牌,得到 %d", got)
|
||||
}
|
||||
if got := n.takeTokens(1, 1, notifyMaxSendsPerChannelPerTick, now.Add(time.Minute)); got != 1 {
|
||||
t.Fatalf("满一个周期应补回 1 个令牌,得到 %d", got)
|
||||
}
|
||||
// 不限流渠道走有限上限,避免单轮被无限积压拖住。
|
||||
if got := n.takeTokens(2, 0, notifyUnlimitedBurstPerTick+10, now); got != notifyUnlimitedBurstPerTick {
|
||||
t.Fatalf("不限流应返回每轮上限 %d,得到 %d", notifyUnlimitedBurstPerTick, got)
|
||||
}
|
||||
// 渠道之间的令牌桶互相独立。
|
||||
if got := n.takeTokens(1, 1, notifyMaxSendsPerChannelPerTick, now.Add(time.Millisecond)); got != 0 {
|
||||
t.Fatalf("渠道 1 的桶应仍然为空,得到 %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyTakeTokensKeepsUnusedTokens 锁住「只取 want 个」的语义。
|
||||
//
|
||||
// 曾经的实现把桶整个抽空后才由调用方截断,于是 rate=100/min 的渠道攒满桶、
|
||||
// 一轮只用 5 条,剩下 95 个令牌直接丢弃;渠道这一轮没有待发投递时同样照扣。
|
||||
// 结果是注释声称的「积压时可以一次性冲 rate_per_min 条」在任何情况下都做不到。
|
||||
func TestNotifyTakeTokensKeepsUnusedTokens(t *testing.T) {
|
||||
n := &Notifier{buckets: map[int64]*notifyBucket{}}
|
||||
now := time.Now()
|
||||
// 桶初始为满(100),本轮只要 5 个。
|
||||
if got := n.takeTokens(1, 100, 5, now); got != 5 {
|
||||
t.Fatalf("want=5 时应恰好取 5 个令牌,得到 %d", got)
|
||||
}
|
||||
// 关键断言:余下的 95 个必须还在桶里,而不是被抽空丢弃。
|
||||
// 不推进时间,确保取到的只可能来自存量而非补充。
|
||||
if got := n.takeTokens(1, 100, 95, now); got != 95 {
|
||||
t.Fatalf("剩余令牌应仍可取用(期望 95),得到 %d——桶被整轮抽空了", got)
|
||||
}
|
||||
if got := n.takeTokens(1, 100, 1, now); got != 0 {
|
||||
t.Fatalf("桶已取尽,应返回 0,得到 %d", got)
|
||||
}
|
||||
// want<=0 不应扣减任何令牌(空轮不收费)。
|
||||
n2 := &Notifier{buckets: map[int64]*notifyBucket{}}
|
||||
if got := n2.takeTokens(1, 20, 0, now); got != 0 {
|
||||
t.Fatalf("want=0 应返回 0,得到 %d", got)
|
||||
}
|
||||
if got := n2.takeTokens(1, 20, 20, now); got != 20 {
|
||||
t.Fatalf("want=0 的那次不该消耗令牌,应仍可取满 20,得到 %d", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDigestTickPlanDecouplesBatchSizeFromSendBudget 钉住汇总模式的两个量纲。
|
||||
//
|
||||
// 汇总批次的大小一旦跟每轮请求预算挂上,rate_per_min=20 的渠道就只能在每个
|
||||
// 3 秒 tick 里补到 1 个令牌,于是每条汇总消息只装 1 个漏洞——功能上等于没有
|
||||
// 汇总,而消息头部还写着「近 30 分钟新增 1 个漏洞」。这个退化不会报错,
|
||||
// 现有的端到端用例也看不出来(它们手动给 stepDigest 传一个够大的 limit,
|
||||
// 绕过了 step 里的额度计算),所以在这里直接断言决策本身。
|
||||
func TestDigestTickPlanDecouplesBatchSizeFromSendBudget(t *testing.T) {
|
||||
tokens, claimLimit := digestTickPlan()
|
||||
// 一批 = 一条消息 = 一次请求 = 一个令牌。令牌的单位是消息,不是漏洞。
|
||||
if tokens != 1 {
|
||||
t.Fatalf("汇总一批只发一条消息,应恰好消耗 1 个令牌,得到 %d", tokens)
|
||||
}
|
||||
if claimLimit != db.MaxDigestBatchSize {
|
||||
t.Fatalf("汇总批次大小应为内存上界 db.MaxDigestBatchSize=%d,得到 %d",
|
||||
db.MaxDigestBatchSize, claimLimit)
|
||||
}
|
||||
// 关键关系:批次大小必须远大于每轮请求预算。两者一旦同量级,
|
||||
// 说明又把「发几条消息」和「一批装几条漏洞」混成了一个数。
|
||||
if claimLimit <= notifyMaxSendsPerChannelPerTick {
|
||||
t.Fatalf("汇总批次大小 %d 不应受每轮请求预算 %d 约束——"+
|
||||
"请求预算是由租约倒推的「发几次请求」,与「一批装几条漏洞」是两个量纲",
|
||||
claimLimit, notifyMaxSendsPerChannelPerTick)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNotifyTickBudgetFitsWithinLease 是又一条防漂移断言。
|
||||
//
|
||||
// 单渠道每轮的投递条数上限(notifyMaxSendsPerChannelPerTick)是从租约时长倒推的:
|
||||
// 一轮里串行投递的最坏耗时必须 < 租约,否则后几条还没发完租约就过期,
|
||||
// 多实例部署时对端会把它们重新领走、重复发送。这三个常量分处不同位置,
|
||||
// 改任意一个都可能打破关系而不会有任何报错——所以在这里钉死。
|
||||
func TestNotifyTickBudgetFitsWithinLease(t *testing.T) {
|
||||
worst := time.Duration(notifyMaxSendsPerChannelPerTick) * notifySendTimeout
|
||||
if worst >= notifyLease {
|
||||
t.Fatalf("单渠道一轮的最坏耗时 %v 不应达到或超过租约 %v"+
|
||||
"(notifyMaxSendsPerChannelPerTick=%d × notifySendTimeout=%v)——"+
|
||||
"改这三个常量中的任意一个都要同步检查另外两个",
|
||||
worst, notifyLease, notifyMaxSendsPerChannelPerTick, notifySendTimeout)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user