Files
artex/db/notification.go
T
dela 0335d572de
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
First Commit
2026-10-09 08:38:16 +08:00

552 lines
22 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package db
import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"log"
"strings"
"time"
"github.com/Autumn-27/artex/notify"
)
// 本文件是 IM 推送的渠道配置与事件层。投递任务的领取与状态流转见
// db/notification_delivery.go。
//
// 两条不变量,改这个文件时务必保持:
//
// 1. 写漏洞的事务(RecordFindingTx)只调用 InsertNotificationEventTx 做一次盲插,
// 不读任何通知相关的表、不做过滤匹配。任何在这里引入的读操作都可能因为
// 用户配错的过滤条件而污染甚至中止漏洞写入事务。
// 2. 过滤匹配永不报错:配置畸形一律按「命中」处理(见 notify.Match)。宁可多推,
// 不可漏推。
// ErrNotificationChannelNotFound 渠道不存在。
var ErrNotificationChannelNotFound = errors.New("通知渠道不存在")
// 投递状态。
const (
NotifyStatePending = "pending" // 待发
NotifyStateSending = "sending" // 已被某个 dispatcher 领取,租约未到期
NotifyStateSent = "sent" // 已送达
NotifyStateFailed = "failed" // 重试耗尽或永久失败,可手动重发
NotifyStateSkipped = "skipped" // 渠道已停用,不再发送
)
// 推送模式。
const (
NotifyModeRealtime = "realtime"
NotifyModeDigest = "digest"
)
// ValidNotifyMode 白名单校验推送模式(与 findings.status 同理:不用 DB CHECK,
// 便于后续扩展)。
func ValidNotifyMode(m string) bool {
return m == NotifyModeRealtime || m == NotifyModeDigest
}
// NotificationChannel 是一个渠道实例配置。Config 与 Filter 保持原始 JSON,
// 解析交给 notify 包——db 层不理解它们的字段含义。
type NotificationChannel struct {
ID int64 `json:"id"`
Name string `json:"name"`
Kind string `json:"kind"`
Mode string `json:"mode"`
Config json.RawMessage `json:"config"`
Filter json.RawMessage `json:"filter"`
// Enabled 用指针是为了区分「没传这个字段」与「显式传 false」——
// 前端开关控件只提交被改动的字段。
Enabled *bool `json:"enabled,omitempty"`
RatePerMin int `json:"rate_per_min"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// IsEnabled 返回渠道是否启用;Enabled 为 nil(未加载)时按启用处理。
func (c *NotificationChannel) IsEnabled() bool { return c.Enabled == nil || *c.Enabled }
// NotificationEvent 是一条事件事实。
type NotificationEvent struct {
ID int64 `json:"id"`
Kind string `json:"kind"`
FindingID int64 `json:"finding_id"`
Snapshot json.RawMessage `json:"snapshot"`
CreatedAt time.Time `json:"created_at"`
}
const notificationChannelCols = `id, name, kind, enabled, config, mode, filter, rate_per_min, created_at, updated_at`
func scanNotificationChannel(sc interface{ Scan(...any) error }) (*NotificationChannel, error) {
var c NotificationChannel
var enabled bool
if err := sc.Scan(&c.ID, &c.Name, &c.Kind, &enabled, &c.Config, &c.Mode, &c.Filter, &c.RatePerMin, &c.CreatedAt, &c.UpdatedAt); err != nil {
return nil, err
}
c.Enabled = &enabled
return &c, nil
}
// ListNotificationChannels 返回全部渠道实例,启用的排在前面、同级按 id。
// 排序放在 SQL 里是为了让 UI 与 dispatcher 看到同一个稳定顺序。
func (d *DB) ListNotificationChannels(ctx context.Context) ([]*NotificationChannel, error) {
rows, err := d.QueryContext(ctx, `SELECT `+notificationChannelCols+` FROM notification_channels
ORDER BY enabled DESC, id`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*NotificationChannel{}
for rows.Next() {
c, err := scanNotificationChannel(rows)
if err != nil {
return nil, err
}
out = append(out, c)
}
return out, rows.Err()
}
// NotificationChannelByID 取单个渠道。
func (d *DB) NotificationChannelByID(ctx context.Context, id int64) (*NotificationChannel, error) {
row := d.QueryRowContext(ctx, `SELECT `+notificationChannelCols+` FROM notification_channels WHERE id=$1`, id)
c, err := scanNotificationChannel(row)
if err == sql.ErrNoRows {
return nil, ErrNotificationChannelNotFound
}
return c, err
}
// SaveNotificationChannel 新建或更新一个渠道。
//
// 更新时只覆盖调用方显式给出的字段(非 nil / 非空),这样前端可以提交局部
// 修改的抽屉表单,而不必回传 config 里那些它没展示的字段——回传反而会造成
// 「掩码值把真实密钥覆盖掉」的事故。
func (d *DB) SaveNotificationChannel(ctx context.Context, c *NotificationChannel) (int64, error) {
if c.Mode == "" {
c.Mode = NotifyModeRealtime
}
// 这里刻意**不**对 0 做任何加工:0 是合法配置,含义是「不限流」。
//
// 曾经写成 `if c.RatePerMin <= 0 { c.RatePerMin = 默认值 }`,本意是「未指定时
// 给个安全默认」,但那把「显式设成 0」也一起吞掉了——文档、UI 提示与
// takeTokens 都把 0 解释为不限流,唯独这里悄悄改成 20(钉钉/企微/Telegram)
// 或 100(飞书),操作者以为放开了限流、实际被 20/分钟卡着且没有任何提示。
//
// 「未指定」与「显式 0」的区别只有调用方知道(请求体里字段缺省 vs 明确传 0),
// 所以默认值由 server 层在字段缺省时填,见 notifyCreateChannel。
if c.RatePerMin < 0 {
return 0, errors.New("限流值不能为负")
}
if c.Config == nil {
c.Config = json.RawMessage(`{}`)
}
if c.Filter == nil {
c.Filter = json.RawMessage(`{}`)
}
enabled := c.IsEnabled()
if c.ID == 0 {
var id int64
err := d.QueryRowContext(ctx, `INSERT INTO notification_channels(name,kind,enabled,config,mode,filter,rate_per_min)
VALUES($1,$2,$3,$4,$5,$6,$7) RETURNING id`,
c.Name, c.Kind, enabled, string(c.Config), c.Mode, string(c.Filter), c.RatePerMin).Scan(&id)
return id, err
}
res, err := d.ExecContext(ctx, `UPDATE notification_channels
SET name=$2, kind=$3, enabled=$4, config=$5, mode=$6, filter=$7, rate_per_min=$8
WHERE id=$1`,
c.ID, c.Name, c.Kind, enabled, string(c.Config), c.Mode, string(c.Filter), c.RatePerMin)
if err != nil {
return 0, err
}
if n, _ := res.RowsAffected(); n == 0 {
return 0, ErrNotificationChannelNotFound
}
return c.ID, nil
}
// SetNotificationChannelEnabled 切换启停。
//
// 停用一个渠道时,把它尚未发出的投递一并标记为 skipped:否则重新启用后
// 会突然收到一批「停用期间积压」的旧漏洞,时效已失且容易误判为新增。
func (d *DB) SetNotificationChannelEnabled(ctx context.Context, id int64, enabled bool) error {
return d.WithEvidenceTx(ctx, func(tx *sql.Tx) error {
res, err := tx.ExecContext(ctx, `UPDATE notification_channels SET enabled=$2 WHERE id=$1`, id, enabled)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotificationChannelNotFound
}
if !enabled {
if _, err := tx.ExecContext(ctx, `UPDATE notification_deliveries SET state=$2, last_error=$3
WHERE channel_id=$1 AND state IN ($4,$5)`,
id, NotifyStateSkipped, "渠道已停用", NotifyStatePending, NotifyStateSending); err != nil {
return err
}
}
return nil
})
}
// DeleteNotificationChannel 删除渠道。其投递历史随外键级联删除
// (渠道配置都没了,历史无从解读)。
func (d *DB) DeleteNotificationChannel(ctx context.Context, id int64) error {
res, err := d.ExecContext(ctx, `DELETE FROM notification_channels WHERE id=$1`, id)
if err != nil {
return err
}
if n, _ := res.RowsAffected(); n == 0 {
return ErrNotificationChannelNotFound
}
return nil
}
// RecordNotificationEventTx 在调用方的事务里**尽力**写入一条推送事件。
//
// 这是漏洞写入路径上唯一的通知相关改动:一次 INSERT,不读任何表、不认识渠道、
// 不跑过滤。事务提交即保证「漏洞落库」与「推送任务存在」原子一致,
// 不存在提交成功却没入队、消息永久丢失的窗口。
//
// 两个关键设计,都不是随手写的:
//
// 1. **为什么用 SAVEPOINT**:PostgreSQL 里事务内任一语句报错会让整个事务进入
// aborted 状态,此后所有语句(含 COMMIT)一律失败。所以「忽略这条 INSERT
// 的错误、让调用方继续提交」在 PG 里是做不到的——除非用保存点把错误隔离在
// 这一条语句上。没有保存点,就只剩「整笔回滚」这一个选项。
//
// 2. **为什么整笔回滚是错的**:推送是便利功能,漏洞记录才是产品本身。一个通知
// 表的问题(旧库未迁移、磁盘瞬时故障)不该让高危漏洞存不进库。所以这里隔离
// 错误、记日志、返回 false,让漏洞写入照常提交——代价是丢掉这一条推送。
// 返回 bool 而非 error 是刻意的:调用方不该把它当作会影响写入成败的错误。
func RecordNotificationEventTx(ctx context.Context, tx *sql.Tx, kind string, findingID int64, snap notify.Snapshot) bool {
raw, err := json.Marshal(snap)
if err != nil {
log.Printf("[notify] 序列化推送事件失败 finding=%d: %v", findingID, err)
return false
}
if _, err := tx.ExecContext(ctx, `SAVEPOINT notify_event`); err != nil {
log.Printf("[notify] 建立保存点失败 finding=%d: %v", findingID, err)
return false
}
if _, err := tx.ExecContext(ctx, `INSERT INTO notification_events(kind,finding_id,snapshot) VALUES($1,$2,$3)`,
kind, findingID, string(raw)); err != nil {
log.Printf("[notify] 写入推送事件失败 finding=%d(漏洞记录不受影响): %v", findingID, err)
// 回滚到保存点,把事务从 aborted 状态里救回来。
if _, rbErr := tx.ExecContext(ctx, `ROLLBACK TO SAVEPOINT notify_event`); rbErr != nil {
log.Printf("[notify] 回滚到保存点失败 finding=%d: %v", findingID, rbErr)
}
return false
}
// 释放保存点,避免长事务里积攒无用的保存点。
_, _ = tx.ExecContext(ctx, `RELEASE SAVEPOINT notify_event`)
return true
}
// AddNotificationEvent 是 InsertNotificationEventTx 的独立事务版本,供不在
// 既有事务里的调用点使用(如渠道的「发送测试消息」,它没有真实 finding)。
func (d *DB) AddNotificationEvent(ctx context.Context, kind string, findingID int64, snap notify.Snapshot) (int64, error) {
raw, err := json.Marshal(snap)
if err != nil {
return 0, fmt.Errorf("序列化通知事件快照失败: %w", err)
}
var id int64
err = d.QueryRowContext(ctx, `INSERT INTO notification_events(kind,finding_id,snapshot) VALUES($1,$2,$3) RETURNING id`,
kind, findingID, string(raw)).Scan(&id)
return id, err
}
// FanOutPendingEvents 把尚未分派的漏洞事件按当前启用的渠道展开成投递任务,
// 返回本轮处理的事件数与新建的投递数。
//
// 整轮操作在一个事务里:事件用 FOR UPDATE SKIP LOCKED 领取,多个进程同时跑
// 也各自领到不同的行(项目里归档队列的领取用的是同一套手法,见
// db/task_archives.go 的 completeNextArchiveJob)。
//
// 过滤匹配刻意放在 Go 侧而非 SQL:渠道的过滤条件是一组可选字段的 JSONB,
// 用 SQL 表达六种组合的匹配会让查询难以维护,而渠道数量是「人手配的几条」,
// 全量加载后在内存里逐条比对更快也更好测。
//
// 未命中任何渠道的事件同样会被标记 fanned_out ——否则它会永远留在待分派集合里,
// 每个 tick 被重扫一遍。
func (d *DB) FanOutPendingEvents(ctx context.Context, limit int) (eventCount, deliveryCount int, err error) {
if limit <= 0 {
limit = 200
}
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return 0, 0, err
}
defer tx.Rollback() //nolint:errcheck // 提交成功后是 no-op
channels, err := listEnabledNotificationChannelsTx(ctx, tx)
if err != nil {
return 0, 0, err
}
rows, err := tx.QueryContext(ctx, `SELECT id, kind, finding_id, snapshot FROM notification_events
WHERE NOT fanned_out ORDER BY id FOR UPDATE SKIP LOCKED LIMIT $1`, limit)
if err != nil {
return 0, 0, err
}
var (
events []NotificationEvent
parsedSnaps []notify.Snapshot
)
for rows.Next() {
var ev NotificationEvent
if err := rows.Scan(&ev.ID, &ev.Kind, &ev.FindingID, &ev.Snapshot); err != nil {
rows.Close()
return 0, 0, err
}
var snap notify.Snapshot
// 快照是我们自己写的,理论上必定可解析;解析失败不阻断投递流程,
// 但这条事件会因字段全空而被所有带过滤条件的渠道跳过——宁可少推一条
// 也不让一个坏行卡死整个队列。
_ = json.Unmarshal(ev.Snapshot, &snap)
// kind 以行内值为准:快照里那份是渲染用的副本,可能被旧版本写过。
snap.Kind = ev.Kind
events = append(events, ev)
parsedSnaps = append(parsedSnaps, snap)
}
rows.Close()
if err := rows.Err(); err != nil {
return 0, 0, err
}
if len(events) == 0 {
return 0, 0, tx.Commit()
}
type pending struct {
eventID int64
channelID int64
}
var toInsert []pending
for i, snap := range parsedSnaps {
for _, ch := range channels {
if !notify.Match(notify.ParseFilter(ch.Filter), snap) {
continue
}
toInsert = append(toInsert, pending{eventID: events[i].ID, channelID: ch.ID})
}
}
if len(toInsert) > 0 {
var (
vals []string
args []any
)
for _, p := range toInsert {
vals = append(vals, fmt.Sprintf("($%d,$%d)", len(args)+1, len(args)+2))
args = append(args, p.eventID, p.channelID)
}
if _, err := tx.ExecContext(ctx, `INSERT INTO notification_deliveries(event_id,channel_id) VALUES `+strings.Join(vals, ","), args...); err != nil {
return 0, 0, err
}
}
// 标记本轮事件已分派。未命中任何渠道的事件也一起标记(见函数注释)。
ids := make([]string, 0, len(events))
markArgs := make([]any, 0, len(events))
for _, ev := range events {
markArgs = append(markArgs, ev.ID)
ids = append(ids, fmt.Sprintf("$%d", len(markArgs)))
}
if _, err := tx.ExecContext(ctx, `UPDATE notification_events SET fanned_out=true WHERE id IN (`+strings.Join(ids, ",")+`)`, markArgs...); err != nil {
return 0, 0, err
}
return len(events), len(toInsert), tx.Commit()
}
// listEnabledNotificationChannelsTx 在事务里取启用中的渠道。数量很少,
// 不做分页也不加缓存——缓存会引入「改了配置何时生效」这个额外的时序问题。
func listEnabledNotificationChannelsTx(ctx context.Context, tx *sql.Tx) ([]*NotificationChannel, error) {
rows, err := tx.QueryContext(ctx, `SELECT id, name, kind, config, mode, filter, rate_per_min
FROM notification_channels WHERE enabled ORDER BY id`)
if err != nil {
return nil, err
}
defer rows.Close()
out := []*NotificationChannel{}
for rows.Next() {
var c NotificationChannel
if err := rows.Scan(&c.ID, &c.Name, &c.Kind, &c.Config, &c.Mode, &c.Filter, &c.RatePerMin); err != nil {
return nil, err
}
out = append(out, &c)
}
return out, rows.Err()
}
// NotificationAssetNames 把资产 id 解析成简短展示名,供推送消息使用。
//
// 返回顺序与入参一致、长度可能小于入参(不存在的 id 被跳过)。保持入参顺序是
// 为了让同一条漏洞的消息在多次投递里资产顺序稳定——否则重试后收到的消息里
// 资产次序变了,会被误读成「资产变了」。
func (d *DB) NotificationAssetNames(ctx context.Context, ids []int64) ([]string, error) {
if len(ids) == 0 {
return nil, nil
}
ph, args := placeholders(1, ids)
rows, err := d.QueryContext(ctx, `SELECT id, type, domain, ip, url, app_name, bundle_id FROM assets WHERE id IN (`+ph+`)`, args...)
if err != nil {
return nil, err
}
defer rows.Close()
labels := map[int64]string{}
for rows.Next() {
var (
id int64
typ string
domain, ip, url sql.NullString
appName, bundleID sql.NullString
)
if err := rows.Scan(&id, &typ, &domain, &ip, &url, &appName, &bundleID); err != nil {
return nil, err
}
labels[id] = assetDisplayName(typ, domain.String, ip.String, url.String, appName.String, bundleID.String)
}
if err := rows.Err(); err != nil {
return nil, err
}
out := make([]string, 0, len(ids))
seen := map[int64]bool{}
for _, id := range ids {
if seen[id] {
continue
}
seen[id] = true
if label, ok := labels[id]; ok && label != "" {
out = append(out, label)
}
}
return out, nil
}
// assetDisplayName 按资产类型挑选最具辨识度的标识。
// 兜底返回空串,由调用方决定怎么呈现「名字取不到的资产」——本函数不臆造占位符,
// 否则「资产#42」这种噪音会混进推送消息里,读者还以为是真实域名。
func assetDisplayName(typ, domain, ip, url, appName, bundleID string) string {
pick := func(vals ...string) string {
for _, v := range vals {
if strings.TrimSpace(v) != "" {
return v
}
}
return ""
}
switch typ {
case "root_domain", "subdomain":
return domain
case "ip":
return ip
case "app":
return pick(appName, bundleID)
case "service", "endpoint":
return pick(url, domain, ip)
default:
return pick(domain, ip, url, appName)
}
}
// SetFindingStatusWithNotify 更新漏洞处置状态,并在同一事务里登记一条状态变更
// 推送事件。
//
// 返回 from=变更前的状态;found=漏洞是否存在;notified=事件是否登记成功。
//
// 三条刻意的行为:
// - 状态未实际变化时不登记事件。前端抽屉重复提交同一个值、或自动化脚本
// 幂等重放,都不该产出推送噪音。
// - 漏洞不存在时返回 found=false 且不做任何写入,由调用方翻译成 404。
// - 事件登记失败不影响状态更新(见 RecordNotificationEventTx 的保存点说明),
// 所以 notified=false 时状态已经改成功了,调用方不应因此报错。
func (d *DB) SetFindingStatusWithNotify(ctx context.Context, id int64, status string) (from string, found bool, notified bool, err error) {
err = d.WithEvidenceTx(ctx, func(tx *sql.Tx) error {
var txErr error
from, found, _, notified, txErr = SetFindingStatusTx(ctx, tx, id, status)
return txErr
})
return from, found, notified, err
}
// SetFindingStatusTx 在**调用方的事务**内更新漏洞状态并登记状态变更推送事件。
//
// 抽成事务级函数是为了让所有改状态的路径共用同一套语义——此前只有
// patchFinding 走带通知的版本,而**复测结论为「已修复」时**(finding_retests
// 里那条 `UPDATE findings SET status=...`)是直接写库的,于是配了
// `on_status_change` 的渠道对这类状态流转完全收不到推送:界面上状态悄悄变了,
// 运维要到打开平台才发现。
//
// 返回 from=变更前状态、found=漏洞是否存在、changed=状态是否真的变了、
// notified=事件是否登记成功(登记失败不影响状态更新,见 RecordNotificationEventTx)。
func SetFindingStatusTx(ctx context.Context, tx *sql.Tx, id int64, status string) (from string, found bool, changed bool, notified bool, err error) {
var (
vulnclass, name, severity, summary string
taskID sql.NullInt64
assetIDs []byte
)
scanErr := tx.QueryRowContext(ctx, `SELECT vulnclass, name, severity, summary, task_id, asset_ids, status
FROM findings WHERE id=$1 FOR UPDATE`, id).
Scan(&vulnclass, &name, &severity, &summary, &taskID, &assetIDs, &from)
if scanErr == sql.ErrNoRows {
return "", false, false, false, nil
}
if scanErr != nil {
return "", false, false, false, scanErr
}
found = true
if from == status {
// 状态没有真的变化就不登记事件:重复提交同一个值、幂等重放都不该
// 产生推送噪音。
return from, true, false, false, nil
}
if _, err := tx.ExecContext(ctx, `UPDATE findings SET status=$2 WHERE id=$1`, id, status); err != nil {
return from, true, false, false, err
}
var assets []int64
_ = json.Unmarshal(assetIDs, &assets)
notified = RecordNotificationEventTx(ctx, tx, notify.EventFindingStatusChanged, id, notify.Snapshot{
Kind: notify.EventFindingStatusChanged,
FindingID: id,
TaskID: taskID.Int64,
VulnClass: vulnclass,
Name: name,
Severity: severity,
Summary: summary,
AssetIDs: assets,
FromStatus: from,
ToStatus: status,
})
return from, true, true, notified, nil
}
// NotificationStats 是通知页顶部的概览计数。
type NotificationStats struct {
Channels int `json:"channels"`
ChannelsOn int `json:"channels_on"`
Pending int `json:"pending"`
Failed int `json:"failed"`
SentToday int `json:"sent_today"`
BacklogAgeMS int64 `json:"backlog_age_ms"` // 最老的待发投递距今毫秒数
}
// NotificationStatsSnapshot 汇总通知系统的健康度。
// BacklogAgeMS 是「推送是不是卡住了」最直接的指标——比 pending 计数有用得多,
// 因为积压 3 条和积压 3 条的差别可以是从 3 秒到 3 小时。
func (d *DB) NotificationStatsSnapshot(ctx context.Context) (*NotificationStats, error) {
var s NotificationStats
if err := d.QueryRowContext(ctx, `SELECT
(SELECT count(*) FROM notification_channels),
(SELECT count(*) FROM notification_channels WHERE enabled),
(SELECT count(*) FROM notification_deliveries WHERE state IN ($1,$2)),
(SELECT count(*) FROM notification_deliveries WHERE state=$3),
(SELECT count(*) FROM notification_deliveries WHERE state=$4 AND sent_at >= date_trunc('day', now())),
COALESCE((SELECT EXTRACT(EPOCH FROM (now() - min(created_at))) * 1000 FROM notification_deliveries WHERE state=$1), 0)::bigint`,
NotifyStatePending, NotifyStateSending, NotifyStateFailed, NotifyStateSent).
Scan(&s.Channels, &s.ChannelsOn, &s.Pending, &s.Failed, &s.SentToday, &s.BacklogAgeMS); err != nil {
return nil, err
}
return &s, nil
}