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