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
detections / detections (push) Waiting to run
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
875 lines
32 KiB
Go
875 lines
32 KiB
Go
package db
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"testing"
|
||
"time"
|
||
|
||
"github.com/Autumn-27/artex/notify"
|
||
)
|
||
|
||
// 本文件的用例都会真连 PostgreSQL(无库时跳过)。这些 SQL 用到了
|
||
// FOR UPDATE SKIP LOCKED、make_interval、JSONB、多行 IN(...) 占位符拼接,
|
||
// 都是「编译通过但可能运行时报错」的写法,必须实跑才算验证过。
|
||
|
||
func notifyTestDB(t *testing.T) *DB {
|
||
t.Helper()
|
||
d, err := Open(testDSN(t))
|
||
if err != nil {
|
||
t.Skipf("postgres unavailable (%v) — skipping", err)
|
||
}
|
||
t.Cleanup(func() { d.Close() })
|
||
return d
|
||
}
|
||
|
||
// newTestChannel 建一个渠道,测试结束自动删除。
|
||
func newTestChannel(t *testing.T, d *DB, kind, mode string, filter string) *NotificationChannel {
|
||
t.Helper()
|
||
if filter == "" {
|
||
filter = `{}`
|
||
}
|
||
ch := &NotificationChannel{
|
||
Name: "测试渠道-" + t.Name(),
|
||
Kind: kind,
|
||
Mode: mode,
|
||
Config: json.RawMessage(`{"webhook":"https://example.com/hook"}`),
|
||
Filter: json.RawMessage(filter),
|
||
RatePerMin: 100,
|
||
}
|
||
id, err := d.SaveNotificationChannel(context.Background(), ch)
|
||
if err != nil {
|
||
t.Fatalf("建渠道失败: %v", err)
|
||
}
|
||
t.Cleanup(func() { d.Exec(`DELETE FROM notification_channels WHERE id=$1`, id) })
|
||
ch.ID = id
|
||
return ch
|
||
}
|
||
|
||
// addTestEvent 直接写一条事件(不经 finding),用于测试分派与投递。
|
||
func addTestEvent(t *testing.T, d *DB, kind string, findingID int64, snap notify.Snapshot) int64 {
|
||
t.Helper()
|
||
snap.Kind = kind
|
||
snap.FindingID = findingID
|
||
id, err := d.AddNotificationEvent(context.Background(), kind, findingID, snap)
|
||
if err != nil {
|
||
t.Fatalf("写事件失败: %v", err)
|
||
}
|
||
t.Cleanup(func() { d.Exec(`DELETE FROM notification_events WHERE id=$1`, id) })
|
||
return id
|
||
}
|
||
|
||
func TestNotificationAssetNamesResolvesAndPreservesOrder(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
// 三类资产各有各的展示口径:域名、IP、URL。
|
||
insertAsset := func(query, value string) int64 {
|
||
t.Helper()
|
||
var id int64
|
||
if err := d.QueryRow(query, value).Scan(&id); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
return id
|
||
}
|
||
domID := insertAsset(`INSERT INTO assets(type, domain) VALUES('subdomain',$1) RETURNING id`, "a.example.com")
|
||
ipID := insertAsset(`INSERT INTO assets(type, ip) VALUES('ip',$1) RETURNING id`, "10.1.2.3")
|
||
svcID := insertAsset(`INSERT INTO assets(type, url) VALUES('service',$1) RETURNING id`, "https://a.example.com/admin")
|
||
t.Cleanup(func() {
|
||
d.Exec(`DELETE FROM assets WHERE id IN ($1,$2,$3)`, domID, ipID, svcID)
|
||
})
|
||
|
||
// 传入顺序刻意乱序,且含一个不存在的 id。
|
||
got, err := d.NotificationAssetNames(ctx, []int64{svcID, 999999999, domID, ipID, svcID})
|
||
if err != nil {
|
||
t.Fatalf("解析资产名失败: %v", err)
|
||
}
|
||
want := []string{"https://a.example.com/admin", "a.example.com", "10.1.2.3"}
|
||
if len(got) != len(want) {
|
||
t.Fatalf("资产名数量不符,期望 %v 得到 %v", want, got)
|
||
}
|
||
for i := range want {
|
||
if got[i] != want[i] {
|
||
t.Fatalf("顺序/取值不符,期望 %v 得到 %v", want, got)
|
||
}
|
||
}
|
||
}
|
||
|
||
// TestRecordNotificationEventTxUnwindsOnFailure 是保存点机制的核心用例:
|
||
// 在事务里先让 notification_events 的写入必然失败(临时加一个恒 false 的约束),
|
||
// 断言 ① 该函数报 false ② 事务没有进入 aborted 状态,后续语句仍能执行。
|
||
//
|
||
// 没有保存点的话,PostgreSQL 会让整个事务作废,后续任何语句都以
|
||
// "current transaction is aborted" 失败——那正是「一个通知表的问题导致
|
||
// 漏洞存不进库」的故障路径。
|
||
//
|
||
// 这里刻意用 **ROLLBACK 收尾而不是 COMMIT**:ALTER TABLE 在 PG 里是事务性的,
|
||
// 一旦提交,那个临时约束就会永久留在 schema 里,把后续所有用例一起打挂。
|
||
// 回滚能自动撤销 DDL,无需手工清理。断言只需要「事务还活着」,
|
||
// 不需要真的提交。
|
||
func TestRecordNotificationEventTxUnwindsOnFailure(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
// 防御性清理:若历史运行留下过这个约束,先摘掉。
|
||
if _, err := d.Exec(`ALTER TABLE notification_events DROP CONSTRAINT IF EXISTS notify_test_never`); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
|
||
tx, err := d.BeginTx(ctx, nil)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
defer tx.Rollback() //nolint:errcheck // 撤销临时约束,见函数注释
|
||
|
||
// NOT VALID:只约束此后写入的行,不去校验库里已有的历史事件
|
||
// (否则存量行违规会导致约束加不上)。
|
||
if _, err := tx.ExecContext(ctx, `ALTER TABLE notification_events ADD CONSTRAINT notify_test_never CHECK (false) NOT VALID`); err != nil {
|
||
t.Fatalf("加临时约束失败: %v", err)
|
||
}
|
||
if RecordNotificationEventTx(ctx, tx, notify.EventFindingCreated, 1, notify.Snapshot{Severity: "high"}) {
|
||
t.Fatal("在必然失败的约束下仍报告写入成功")
|
||
}
|
||
// 关键断言:事务还能用。
|
||
var one int
|
||
if err := tx.QueryRowContext(ctx, `SELECT 1`).Scan(&one); err != nil {
|
||
t.Fatalf("事务已被污染(保存点未生效): %v", err)
|
||
}
|
||
if err := tx.Rollback(); err != nil {
|
||
t.Fatalf("回滚失败: %v", err)
|
||
}
|
||
// 确认 DDL 已随回滚撤销,不给后续用例留雷。
|
||
var exists bool
|
||
if err := d.QueryRow(`SELECT EXISTS(SELECT 1 FROM pg_constraint WHERE conname='notify_test_never')`).Scan(&exists); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if exists {
|
||
t.Fatal("临时约束未被回滚撤销,会污染后续用例")
|
||
}
|
||
}
|
||
|
||
func TestFanOutRoutesEventsByFilter(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
all := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
|
||
onlyCritical := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{"min_severity":"critical"}`)
|
||
sqlOnly := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{"vulnclass_include":["SQL"]}`)
|
||
|
||
highSQL := addTestEvent(t, d, notify.EventFindingCreated, 1001, notify.Snapshot{Severity: "high", VulnClass: "SQL注入"})
|
||
lowXSS := addTestEvent(t, d, notify.EventFindingCreated, 1002, notify.Snapshot{Severity: "low", VulnClass: "XSS"})
|
||
criticalXSS := addTestEvent(t, d, notify.EventFindingCreated, 1003, notify.Snapshot{Severity: "critical", VulnClass: "XSS"})
|
||
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatalf("分派失败: %v", err)
|
||
}
|
||
|
||
cases := []struct {
|
||
name string
|
||
eventID int64
|
||
channel int64
|
||
want bool
|
||
}{
|
||
{"全收渠道收到 high", highSQL, all.ID, true},
|
||
{"全收渠道收到 low", lowXSS, all.ID, true},
|
||
{"仅严重渠道跳过 high", highSQL, onlyCritical.ID, false},
|
||
{"仅严重渠道收到 critical", criticalXSS, onlyCritical.ID, true},
|
||
{"仅SQL渠道收到 SQL", highSQL, sqlOnly.ID, true},
|
||
{"仅SQL渠道跳过 XSS", lowXSS, sqlOnly.ID, false},
|
||
}
|
||
for _, tc := range cases {
|
||
t.Run(tc.name, func(t *testing.T) {
|
||
var exists bool
|
||
if err := d.QueryRowContext(ctx, `SELECT EXISTS(SELECT 1 FROM notification_deliveries WHERE event_id=$1 AND channel_id=$2)`,
|
||
tc.eventID, tc.channel).Scan(&exists); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if exists != tc.want {
|
||
t.Fatalf("投递是否存在: 期望 %v 得到 %v", tc.want, exists)
|
||
}
|
||
})
|
||
}
|
||
|
||
// 再分派一次不应产生重复投递(fanned_out 幂等)。
|
||
events, deliveries, err := d.FanOutPendingEvents(ctx, 100)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if events != 0 || deliveries != 0 {
|
||
t.Fatalf("已分派的事件不应被再次处理,得到 events=%d deliveries=%d", events, deliveries)
|
||
}
|
||
}
|
||
|
||
// TestFanOutMarksEventsWithNoMatchingChannel 覆盖「事件没命中任何渠道」的情况。
|
||
// 这类事件必须照样被标记为已分派,否则它会永远留在待分派集合里、每个 tick 重扫。
|
||
func TestFanOutMarksEventsWithNoMatchingChannel(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
pick := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{"vulnclass_include":["绝不匹配的类型"]}`)
|
||
_ = pick
|
||
|
||
ev := addTestEvent(t, d, notify.EventFindingCreated, 2001, notify.Snapshot{Severity: "high", VulnClass: "XSS"})
|
||
_, deliveries, err := d.FanOutPendingEvents(ctx, 100)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if deliveries != 0 {
|
||
t.Fatalf("不该产生投递,得到 %d", deliveries)
|
||
}
|
||
var fanned bool
|
||
if err := d.QueryRowContext(ctx, `SELECT fanned_out FROM notification_events WHERE id=$1`, ev).Scan(&fanned); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !fanned {
|
||
t.Fatal("未命中渠道的事件也必须标记为已分派,否则会被无限重扫")
|
||
}
|
||
}
|
||
|
||
func TestClaimRealtimeDeliveriesHonorsLeaseAndMode(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
realtime := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
|
||
digest := newTestChannel(t, d, notify.KindDingTalk, NotifyModeDigest, `{}`)
|
||
|
||
addTestEvent(t, d, notify.EventFindingCreated, 3001, notify.Snapshot{Severity: "high", VulnClass: "XSS"})
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
|
||
// 实时领取只应拿到 realtime 渠道的那条,不该动 digest 渠道的。
|
||
got, err := d.ClaimRealtimeDeliveries(ctx, realtime.ID, 10, time.Minute)
|
||
if err != nil {
|
||
t.Fatalf("领取失败: %v", err)
|
||
}
|
||
if len(got) != 1 {
|
||
t.Fatalf("应领到 1 条,得到 %d", len(got))
|
||
}
|
||
if got[0].State != NotifyStateSending || got[0].Attempts != 1 {
|
||
t.Fatalf("领取后应为 sending 且 attempts=1,得到 state=%s attempts=%d", got[0].State, got[0].Attempts)
|
||
}
|
||
// 关联加载的渲染上下文必须齐全(渠道配置 + 事件快照 + finding id)。
|
||
if got[0].Channel == nil || len(got[0].Channel.Config) == 0 {
|
||
t.Fatal("领取结果缺少渠道配置,渲染会失败")
|
||
}
|
||
if got[0].FindingID != 3001 {
|
||
t.Fatalf("finding id 未从事件带出,得到 %d", got[0].FindingID)
|
||
}
|
||
|
||
// 租约未到期,第二次领取应为空——这是「同一行不会被两个 dispatcher 同时投递」
|
||
// 的保证。
|
||
again, err := d.ClaimRealtimeDeliveries(ctx, realtime.ID, 10, time.Minute)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if len(again) != 0 {
|
||
t.Fatalf("租约期内不应重复领取,得到 %d 条", len(again))
|
||
}
|
||
|
||
// digest 渠道的投递不应被实时领取碰到。
|
||
left, err := d.ClaimRealtimeDeliveries(ctx, digest.ID, 10, time.Minute)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if len(left) != 0 {
|
||
t.Fatalf("实时领取不应拿到 digest 渠道的投递,得到 %d 条", len(left))
|
||
}
|
||
}
|
||
|
||
// TestClaimExpiredLeaseRecovers 覆盖崩溃自愈:进程在投递途中挂掉会留下 sending
|
||
// 行,租约到期后必须能被重新领起来,否则这条投递永远卡住。
|
||
func TestClaimExpiredLeaseRecovers(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
|
||
addTestEvent(t, d, notify.EventFindingCreated, 4001, notify.Snapshot{Severity: "high"})
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
first, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute)
|
||
if err != nil || len(first) != 1 {
|
||
t.Fatalf("首次领取失败: %v (%d 条)", err, len(first))
|
||
}
|
||
// 把租约手动推到过去,模拟「租约已过期」。
|
||
if _, err := d.Exec(`UPDATE notification_deliveries SET next_attempt_at = now() - interval '1 minute' WHERE id=$1`, first[0].ID); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
second, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if len(second) != 1 {
|
||
t.Fatalf("租约过期的 sending 行应可被重新领取,得到 %d 条", len(second))
|
||
}
|
||
if second[0].Attempts != 2 {
|
||
t.Fatalf("重新领取应累加尝试次数,得到 %d", second[0].Attempts)
|
||
}
|
||
}
|
||
|
||
func TestClaimSkipsDisabledChannel(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
|
||
addTestEvent(t, d, notify.EventFindingCreated, 5001, notify.Snapshot{Severity: "high"})
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// 停用会把存量待发投递一起标记为 skipped。
|
||
if err := d.SetNotificationChannelEnabled(ctx, ch.ID, false); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
var state string
|
||
if err := d.QueryRow(`SELECT state FROM notification_deliveries WHERE channel_id=$1`, ch.ID).Scan(&state); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if state != NotifyStateSkipped {
|
||
t.Fatalf("停用渠道的存量待发投递应被标记为 skipped,得到 %s", state)
|
||
}
|
||
got, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if len(got) != 0 {
|
||
t.Fatalf("已停用渠道不应能被领取,得到 %d 条", len(got))
|
||
}
|
||
}
|
||
|
||
func TestDigestBatchDueAndStableBatchID(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeDigest, `{}`)
|
||
for i := 0; i < 3; i++ {
|
||
addTestEvent(t, d, notify.EventFindingCreated, int64(6000+i), notify.Snapshot{Severity: "high"})
|
||
}
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
|
||
// 批次刚建、年龄为 0,30 分钟的周期下不该到期。
|
||
due, err := d.DigestBatchDue(ctx, ch.ID, 30*time.Minute)
|
||
if err != nil {
|
||
t.Fatalf("判断批次到期失败: %v", err)
|
||
}
|
||
if due {
|
||
t.Fatal("刚建立的批次不应立即到期")
|
||
}
|
||
|
||
// 把三条投递的创建时间一起推老,模拟一个攒够周期的批次。
|
||
if _, err := d.Exec(`UPDATE notification_deliveries SET created_at = now() - interval '40 minutes' WHERE channel_id=$1`, ch.ID); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
due, err = d.DigestBatchDue(ctx, ch.ID, 30*time.Minute)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !due {
|
||
t.Fatal("超过周期的批次应判定为到期")
|
||
}
|
||
|
||
batch, err := d.ClaimDigestBatch(ctx, ch.ID, MaxDigestBatchSize, time.Minute)
|
||
if err != nil {
|
||
t.Fatalf("领取汇总批次失败: %v", err)
|
||
}
|
||
if len(batch) != 3 {
|
||
t.Fatalf("汇总应一次领走全部 3 条,得到 %d 条", len(batch))
|
||
}
|
||
if batch[0].BatchID == nil {
|
||
t.Fatal("汇总批次必须写 batch_id,否则历史里看不出这几条是一起发的")
|
||
}
|
||
firstBatchID := *batch[0].BatchID
|
||
for _, dl := range batch {
|
||
if dl.BatchID == nil || *dl.BatchID != firstBatchID {
|
||
t.Fatalf("同一批次应共享 batch_id,得到 %v vs %d", dl.BatchID, firstBatchID)
|
||
}
|
||
}
|
||
|
||
// 让这批**整体**失败重排后再领,batch_id 必须保持原值(COALESCE 的作用):
|
||
// 否则一次重试就把「这批是一起发的」这个事实抹掉了。
|
||
//
|
||
// 必须整批重排而不是只重排一条——投递引擎发汇总消息时就是这样处理的
|
||
// (一条消息代表整批,成败与共)。只重排一条的话,其余仍在租约期内,
|
||
// 重领自然只拿到那一条。
|
||
allIDs := make([]int64, 0, len(batch))
|
||
for _, dl := range batch {
|
||
allIDs = append(allIDs, dl.ID)
|
||
}
|
||
if err := d.RescheduleDeliveries(ctx, allIDs, time.Second, "模拟失败"); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// 把租约推到过去,模拟退避时间已到。
|
||
if _, err := d.Exec(`UPDATE notification_deliveries SET next_attempt_at = now() - interval '1 minute' WHERE channel_id=$1`, ch.ID); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
reclaimed, err := d.ClaimDigestBatch(ctx, ch.ID, MaxDigestBatchSize, time.Minute)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if len(reclaimed) != 3 {
|
||
t.Fatalf("重领应拿到全部 3 条,得到 %d", len(reclaimed))
|
||
}
|
||
if reclaimed[0].BatchID == nil || *reclaimed[0].BatchID != firstBatchID {
|
||
t.Fatalf("重试后 batch_id 应保持原值 %d,得到 %v", firstBatchID, reclaimed[0].BatchID)
|
||
}
|
||
}
|
||
|
||
func TestDeliveryStateTransitions(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
|
||
addTestEvent(t, d, notify.EventFindingCreated, 7001, notify.Snapshot{Severity: "high"})
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
got, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute)
|
||
if err != nil || len(got) != 1 {
|
||
t.Fatalf("领取失败: %v (%d)", err, len(got))
|
||
}
|
||
id := got[0].ID
|
||
|
||
if err := d.RescheduleDeliveries(ctx, []int64{id}, time.Second, "网络抖动"); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
var state, lastErr string
|
||
if err := d.QueryRow(`SELECT state, last_error FROM notification_deliveries WHERE id=$1`, id).Scan(&state, &lastErr); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if state != NotifyStatePending || lastErr != "网络抖动" {
|
||
t.Fatalf("重排后应为 pending 并记录原因,得到 state=%s err=%q", state, lastErr)
|
||
}
|
||
|
||
if err := d.FailDeliveries(ctx, []int64{id}, "重试耗尽"); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := d.QueryRow(`SELECT state FROM notification_deliveries WHERE id=$1`, id).Scan(&state); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if state != NotifyStateFailed {
|
||
t.Fatalf("应为 failed,得到 %s", state)
|
||
}
|
||
|
||
// 手动重发要清零重试计数并立即到期,否则会继承旧的失败预算。
|
||
if err := d.RetryNotificationDelivery(ctx, id); err != nil {
|
||
t.Fatalf("重发失败: %v", err)
|
||
}
|
||
var attempts int
|
||
var next time.Time
|
||
if err := d.QueryRow(`SELECT state, attempts, next_attempt_at FROM notification_deliveries WHERE id=$1`, id).Scan(&state, &attempts, &next); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if state != NotifyStatePending || attempts != 0 {
|
||
t.Fatalf("重发后应为 pending 且 attempts=0,得到 state=%s attempts=%d", state, attempts)
|
||
}
|
||
if next.After(time.Now().Add(time.Second)) {
|
||
t.Fatal("重发应立即可领")
|
||
}
|
||
|
||
// 已送达的投递不应能被重发。
|
||
if err := d.MarkDeliveriesSent(ctx, []int64{id}); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := d.RetryNotificationDelivery(ctx, id); err == nil {
|
||
t.Fatal("已送达的投递不该允许重发")
|
||
}
|
||
}
|
||
|
||
func TestListNotificationDeliveriesPagingAndFilter(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
|
||
for i := 0; i < 5; i++ {
|
||
addTestEvent(t, d, notify.EventFindingCreated, int64(8000+i), notify.Snapshot{Severity: "high", Name: "分页测试"})
|
||
}
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if _, err := d.ClaimRealtimeDeliveries(ctx, ch.ID, 10, time.Minute); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
|
||
page1, total, err := d.ListNotificationDeliveries(ctx, NotificationDeliveryFilter{ChannelID: ch.ID, State: NotifyStateSending}, 1, 2)
|
||
if err != nil {
|
||
t.Fatalf("查询失败: %v", err)
|
||
}
|
||
if total != 5 {
|
||
t.Fatalf("总数应为 5,得到 %d", total)
|
||
}
|
||
if len(page1) != 2 {
|
||
t.Fatalf("每页 2 条,得到 %d", len(page1))
|
||
}
|
||
// 新的在前:第一页首条 id 应大于第二页首条。
|
||
page2, _, err := d.ListNotificationDeliveries(ctx, NotificationDeliveryFilter{ChannelID: ch.ID, State: NotifyStateSending}, 2, 2)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if len(page2) != 2 || page2[0].ID >= page1[0].ID {
|
||
t.Fatalf("分页顺序应为新的在前,得到 page1[0]=%d page2[0]=%d", page1[0].ID, page2[0].ID)
|
||
}
|
||
// 渲染上下文必须随历史一起返回,否则列表无法展示「推的是什么」。
|
||
if page1[0].ChannelName == "" || page1[0].FindingID == 0 {
|
||
t.Fatalf("历史项缺少展示字段: %+v", page1[0])
|
||
}
|
||
|
||
// 按状态过滤:没有 pending 的。
|
||
pending, totalPending, err := d.ListNotificationDeliveries(ctx, NotificationDeliveryFilter{ChannelID: ch.ID, State: NotifyStatePending}, 1, 50)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if totalPending != 0 || len(pending) != 0 {
|
||
t.Fatalf("不该有 pending 投递,得到 %d 条 (total=%d)", len(pending), totalPending)
|
||
}
|
||
}
|
||
|
||
func TestSetFindingStatusWithNotifyOnlyEmitsOnRealChange(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
tk, err := d.CreateTask("通知状态变更测试", "目标", nil, 0, 0)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
defer d.DeleteTask(tk.ID)
|
||
es := d.Exploration(tk.ExplorationID)
|
||
f, err := es.RecordFinding(ctx, RecordFindingInput{
|
||
TaskID: tk.ID, Worker: "test", VulnClass: "SQL注入", Name: "状态变更用例",
|
||
Severity: "high", Summary: "摘要",
|
||
})
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
t.Cleanup(func() { d.Exec(`DELETE FROM notification_events WHERE finding_id=$1`, f.FindingID) })
|
||
|
||
// 落库时已登记一条 finding_created 事件,先把它数出来作为基线。
|
||
var base int
|
||
if err := d.QueryRow(`SELECT count(*) FROM notification_events WHERE finding_id=$1`, f.FindingID).Scan(&base); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if base < 1 {
|
||
t.Fatal("漏洞落库应在同一事务里登记一条推送事件")
|
||
}
|
||
|
||
// 改成同一个状态:不该产生事件(避免重复提交刷出推送噪音)。
|
||
from, found, notified, err := d.SetFindingStatusWithNotify(ctx, f.FindingID, "pending")
|
||
if err != nil || !found {
|
||
t.Fatalf("状态设置失败: found=%v err=%v", found, err)
|
||
}
|
||
if notified {
|
||
t.Fatal("状态未变化时不应登记推送事件")
|
||
}
|
||
if from != "pending" {
|
||
t.Fatalf("应返回变更前状态 pending,得到 %q", from)
|
||
}
|
||
|
||
// 真正变更:应登记事件并记录 from/to。
|
||
from, found, notified, err = d.SetFindingStatusWithNotify(ctx, f.FindingID, "fixed")
|
||
if err != nil || !found {
|
||
t.Fatalf("状态设置失败: found=%v err=%v", found, err)
|
||
}
|
||
if !notified {
|
||
t.Fatal("状态实际变更时应登记推送事件")
|
||
}
|
||
if from != "pending" {
|
||
t.Fatalf("from 应为 pending,得到 %q", from)
|
||
}
|
||
var snapshot []byte
|
||
if err := d.QueryRow(`SELECT snapshot FROM notification_events WHERE finding_id=$1 AND kind=$2`,
|
||
f.FindingID, notify.EventFindingStatusChanged).Scan(&snapshot); err != nil {
|
||
t.Fatalf("未找到状态变更事件: %v", err)
|
||
}
|
||
var snap notify.Snapshot
|
||
if err := json.Unmarshal(snapshot, &snap); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if snap.FromStatus != "pending" || snap.ToStatus != "fixed" {
|
||
t.Fatalf("快照里的状态流转不对: %s → %s", snap.FromStatus, snap.ToStatus)
|
||
}
|
||
// 快照要带上渲染所需字段,否则状态变更消息会是空壳。
|
||
if snap.VulnClass != "SQL注入" || snap.Severity != "high" || snap.Name != "状态变更用例" {
|
||
t.Fatalf("快照缺少渲染字段: %+v", snap)
|
||
}
|
||
var status string
|
||
if err := d.QueryRow(`SELECT status FROM findings WHERE id=$1`, f.FindingID).Scan(&status); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if status != "fixed" {
|
||
t.Fatalf("状态应已更新为 fixed,得到 %s", status)
|
||
}
|
||
|
||
// 不存在的漏洞:found=false,不报错。
|
||
if _, found, _, err := d.SetFindingStatusWithNotify(ctx, 999999999, "fixed"); err != nil || found {
|
||
t.Fatalf("不存在的漏洞应返回 found=false 且无错,得到 found=%v err=%v", found, err)
|
||
}
|
||
}
|
||
|
||
func TestNotificationStatsSnapshot(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
|
||
addTestEvent(t, d, notify.EventFindingCreated, 9001, notify.Snapshot{Severity: "high"})
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
stats, err := d.NotificationStatsSnapshot(ctx)
|
||
if err != nil {
|
||
t.Fatalf("统计失败: %v", err)
|
||
}
|
||
if stats.Channels < 1 || stats.ChannelsOn < 1 {
|
||
t.Fatalf("渠道计数不对: %+v", stats)
|
||
}
|
||
if stats.Pending < 1 {
|
||
t.Fatalf("应统计到待发投递: %+v", stats)
|
||
}
|
||
// 刚建的投递积压年龄应接近 0,而不是负数或巨大值。
|
||
if stats.BacklogAgeMS < 0 || stats.BacklogAgeMS > int64(time.Hour/time.Millisecond) {
|
||
t.Fatalf("积压年龄不合法: %d ms", stats.BacklogAgeMS)
|
||
}
|
||
_ = ch
|
||
}
|
||
|
||
func TestNotificationChannelCRUDRoundTrip(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
ch := &NotificationChannel{
|
||
Name: "CRUD 往返",
|
||
Kind: notify.KindEmail,
|
||
Mode: NotifyModeDigest,
|
||
Config: json.RawMessage(`{"host":"smtp.example.com","port":587,"from":"a@b.c","to":["x@y.z"]}`),
|
||
Filter: json.RawMessage(`{"min_severity":"medium","on_status_change":true}`),
|
||
RatePerMin: 42,
|
||
}
|
||
id, err := d.SaveNotificationChannel(ctx, ch)
|
||
if err != nil {
|
||
t.Fatalf("新建失败: %v", err)
|
||
}
|
||
t.Cleanup(func() { d.Exec(`DELETE FROM notification_channels WHERE id=$1`, id) })
|
||
|
||
got, err := d.NotificationChannelByID(ctx, id)
|
||
if err != nil {
|
||
t.Fatalf("读取失败: %v", err)
|
||
}
|
||
if got.Mode != NotifyModeDigest || got.RatePerMin != 42 || got.Name != "CRUD 往返" {
|
||
t.Fatalf("往返字段不一致: %+v", got)
|
||
}
|
||
if !got.IsEnabled() {
|
||
t.Fatal("默认应为启用")
|
||
}
|
||
var cfg map[string]any
|
||
if err := json.Unmarshal(got.Config, &cfg); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if cfg["host"] != "smtp.example.com" {
|
||
t.Fatalf("配置未正确落库: %v", cfg)
|
||
}
|
||
var filter notify.Filter
|
||
if err := json.Unmarshal(got.Filter, &filter); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if filter.MinSeverity != "medium" || !filter.OnStatusChange {
|
||
t.Fatalf("过滤条件未正确落库: %+v", filter)
|
||
}
|
||
|
||
// 更新后再读。
|
||
got.Name = "改名了"
|
||
off := false
|
||
got.Enabled = &off
|
||
if _, err := d.SaveNotificationChannel(ctx, got); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
after, err := d.NotificationChannelByID(ctx, id)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if after.Name != "改名了" || after.IsEnabled() {
|
||
t.Fatalf("更新未生效: %+v", after)
|
||
}
|
||
|
||
// 删除后应报「不存在」而不是静默成功。
|
||
if err := d.DeleteNotificationChannel(ctx, id); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if _, err := d.NotificationChannelByID(ctx, id); err != ErrNotificationChannelNotFound {
|
||
t.Fatalf("期望 ErrNotificationChannelNotFound,得到 %v", err)
|
||
}
|
||
if err := d.DeleteNotificationChannel(ctx, id); err != ErrNotificationChannelNotFound {
|
||
t.Fatalf("重复删除应报不存在,得到 %v", err)
|
||
}
|
||
}
|
||
|
||
// TestSaveNotificationChannelKeepsExplicitZeroRate 锁住一个曾经写错的地方:
|
||
// **0 是合法配置,含义是「不限流」,不能被 db 层当成「未指定」覆盖成默认值**。
|
||
//
|
||
// 历史 bug:SaveNotificationChannel 里写了 `if RatePerMin <= 0 { 取默认值 }`,
|
||
// 于是文档、UI 提示、takeTokens 都按「0=不限流」解释,唯独写库这一层悄悄改成
|
||
// 20(钉钉/企微/Telegram)或 100(飞书)——操作者以为放开了限流、实际被卡着,
|
||
// 而且没有任何提示。「未指定」与「显式 0」的区别只有请求体能表达,
|
||
// 所以默认值在 server 层填(见 notifyCreateChannel),db 层只管存。
|
||
func TestSaveNotificationChannelKeepsExplicitZeroRate(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
// 显式 0(不限流):必须原样存下来。
|
||
unlimited := &NotificationChannel{
|
||
Name: "不限流", Kind: notify.KindDingTalk, RatePerMin: 0,
|
||
Config: json.RawMessage(`{"webhook":"https://example.com/h"}`),
|
||
}
|
||
id, err := d.SaveNotificationChannel(ctx, unlimited)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
t.Cleanup(func() { d.Exec(`DELETE FROM notification_channels WHERE id=$1`, id) })
|
||
got, err := d.NotificationChannelByID(ctx, id)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if got.RatePerMin != 0 {
|
||
t.Fatalf("显式 0 表示不限流,必须原样保存,得到 %d", got.RatePerMin)
|
||
}
|
||
if got.Mode != NotifyModeRealtime {
|
||
t.Fatalf("默认模式应为 realtime,得到 %s", got.Mode)
|
||
}
|
||
|
||
// 负值是非法输入,应被拒绝而不是悄悄改成别的值。
|
||
bad := &NotificationChannel{
|
||
Name: "负限流", Kind: notify.KindDingTalk, RatePerMin: -1,
|
||
Config: json.RawMessage(`{"webhook":"https://example.com/h"}`),
|
||
}
|
||
if _, err := d.SaveNotificationChannel(ctx, bad); err == nil {
|
||
t.Fatal("负限流应被拒绝")
|
||
}
|
||
}
|
||
|
||
// TestDeleteChannelCascadesDeliveries 锁住外键行为:渠道删除后其投递历史一并消失
|
||
// (配置都没了,历史无从解读),但事件本身要留下——它可能还被别的渠道引用。
|
||
func TestDeleteChannelCascadesDeliveries(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeRealtime, `{}`)
|
||
ev := addTestEvent(t, d, notify.EventFindingCreated, 9101, notify.Snapshot{Severity: "high"})
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
var before int
|
||
if err := d.QueryRow(`SELECT count(*) FROM notification_deliveries WHERE channel_id=$1`, ch.ID).Scan(&before); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if before == 0 {
|
||
t.Fatal("前置条件不成立:未产生投递")
|
||
}
|
||
if err := d.DeleteNotificationChannel(ctx, ch.ID); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
var after int
|
||
if err := d.QueryRow(`SELECT count(*) FROM notification_deliveries WHERE channel_id=$1`, ch.ID).Scan(&after); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if after != 0 {
|
||
t.Fatalf("渠道删除后其投递应级联删除,仍有 %d 条", after)
|
||
}
|
||
var evExists bool
|
||
if err := d.QueryRow(`SELECT EXISTS(SELECT 1 FROM notification_events WHERE id=$1)`, ev).Scan(&evExists); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if !evExists {
|
||
t.Fatal("删渠道不应连带删除事件本身")
|
||
}
|
||
}
|
||
|
||
// TestClaimDigestBatchHonorsCallerLimit 覆盖审计指出的一处口子:
|
||
// 汇总渠道此前完全绕过令牌桶——allow 被 takeTokens 扣掉却没人用,
|
||
// rate_per_min 对 digest 模式毫无作用。现在 limit 也参与约束。
|
||
func TestClaimDigestBatchHonorsCallerLimit(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
ch := newTestChannel(t, d, notify.KindDingTalk, NotifyModeDigest, `{}`)
|
||
for i := 0; i < 10; i++ {
|
||
addTestEvent(t, d, notify.EventFindingCreated, int64(7000+i), notify.Snapshot{Severity: "high"})
|
||
}
|
||
if _, _, err := d.FanOutPendingEvents(ctx, 100); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
// 取 limit=3:只能领到 3 条,其余留在库里。
|
||
got, err := d.ClaimDigestBatch(ctx, ch.ID, 3, time.Minute)
|
||
if err != nil {
|
||
t.Fatalf("领取失败: %v", err)
|
||
}
|
||
if len(got) != 3 {
|
||
t.Fatalf("应按调用方限流额度只领 3 条,得到 %d", len(got))
|
||
}
|
||
// limit=0 表示本轮额度用尽:一条都不该领,也不该报错。
|
||
if got, err := d.ClaimDigestBatch(ctx, ch.ID, 0, time.Minute); err != nil || len(got) != 0 {
|
||
t.Fatalf("额度为 0 时应领 0 条且不报错,得到 %d 条 err=%v", len(got), err)
|
||
}
|
||
}
|
||
|
||
// TestFinishFindingRetestEmitsStatusChange 覆盖审计指出的一处完整性缺口:
|
||
// 复测结论为「已修复」时,状态确实变了,但那条 UPDATE 是直接写库的、
|
||
// 绕过了带通知的版本——于是配了 on_status_change 的渠道对这种状态流转
|
||
// 完全收不到推送,界面上状态悄悄变了,运维要打开平台才知道。
|
||
//
|
||
// 这条用例锁住「所有改状态的路径都要登记状态变更事件」。
|
||
func TestFinishFindingRetestEmitsStatusChange(t *testing.T) {
|
||
d := notifyTestDB(t)
|
||
ctx := context.Background()
|
||
|
||
tk, err := d.CreateTask("复测推送测试", "目标", nil, 0, 0)
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
defer d.DeleteTask(tk.ID)
|
||
es := d.Exploration(tk.ExplorationID)
|
||
f, err := es.RecordFinding(ctx, RecordFindingInput{
|
||
TaskID: tk.ID, Worker: "test", VulnClass: "SQL注入", Name: "复测目标",
|
||
Severity: "high", Summary: "摘要",
|
||
})
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
t.Cleanup(func() { d.Exec(`DELETE FROM notification_events WHERE finding_id=$1`, f.FindingID) })
|
||
|
||
// 建一条复测记录并直接推到完成态。
|
||
rt, _, _, err := d.CreateFindingRetest(ctx, f.FindingID, "复核")
|
||
if err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if rt.ConversationID == nil {
|
||
t.Fatal("复测应关联一个会话")
|
||
}
|
||
// 复测必须先进入 running 才能落结论(与真实流程一致)。
|
||
if ok, err := d.StartFindingRetest(ctx, rt.ID); err != nil || !ok {
|
||
t.Fatalf("启动复测失败: ok=%v err=%v", ok, err)
|
||
}
|
||
if err := d.RecordFindingRetestResult(ctx, *rt.ConversationID, "fixed", "已修复", "证据"); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if err := d.FinishFindingRetest(rt.ID, "completed", ""); err != nil {
|
||
t.Fatalf("结束复测失败: %v", err)
|
||
}
|
||
|
||
var status string
|
||
if err := d.QueryRow(`SELECT status FROM findings WHERE id=$1`, f.FindingID).Scan(&status); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if status != FindingFixed {
|
||
t.Fatalf("复测判已修复后状态应为 fixed,得到 %s", status)
|
||
}
|
||
|
||
// 关键断言:必须有一条状态变更事件,且 from/to 正确。
|
||
var snapshot []byte
|
||
err = d.QueryRow(`SELECT snapshot FROM notification_events WHERE finding_id=$1 AND kind=$2 ORDER BY id DESC LIMIT 1`,
|
||
f.FindingID, notify.EventFindingStatusChanged).Scan(&snapshot)
|
||
if err != nil {
|
||
t.Fatalf("复测判已修复应登记状态变更推送事件(否则配了 on_status_change 的渠道收不到): %v", err)
|
||
}
|
||
var snap notify.Snapshot
|
||
if err := json.Unmarshal(snapshot, &snap); err != nil {
|
||
t.Fatal(err)
|
||
}
|
||
if snap.FromStatus != "pending" || snap.ToStatus != FindingFixed {
|
||
t.Fatalf("快照的状态流转不对: %s → %s", snap.FromStatus, snap.ToStatus)
|
||
}
|
||
// 快照要带渲染所需字段,否则推送出来是空壳。
|
||
if snap.Name != "复测目标" || snap.Severity != "high" {
|
||
t.Fatalf("快照缺少渲染字段: %+v", snap)
|
||
}
|
||
}
|